diff --git a/.circleci/config.yml b/.circleci/config.yml index 370424dca86..d9c85cfa042 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -323,6 +323,7 @@ jobs: CHOCOLATEY_CONFIRM_ALL: "true" - run: name: Install Dependencies + no_output_timeout: 30m environment: UV_HTTP_TIMEOUT: "300" command: | @@ -381,6 +382,7 @@ jobs: uv run --no-sync python -m pytest tests/windows_tests/ -v - run: name: Guard against MAX_PATH-busting packaged wheel paths + no_output_timeout: 30m environment: UV_HTTP_TIMEOUT: "300" command: | @@ -3327,22 +3329,12 @@ workflows: matrix: parameters: suite: [management, accounting, database, providers, extensions, mcp, sdk, cost, browser] - filters: - branches: - only: - - main - - /litellm_.*/ - integration_contracts: name: integration-<< matrix.suite >>-replica matrix: parameters: suite: [management, database] mode: [replica] - filters: - branches: - only: - - main - - /litellm_.*/ build_and_test: unless: or: @@ -3350,101 +3342,60 @@ workflows: - not: equal: ["", << pipeline.parameters.routing_parity_base >>] jobs: - - using_litellm_on_windows: - filters: &main_branches - branches: - only: - - main - - /litellm_.*/ - - unit: - filters: *main_branches + - using_litellm_on_windows + - unit - provider_replay_harness - - base_sdk_install: - filters: *main_branches - - local_testing_part1: - filters: *main_branches - - local_testing_part2: - filters: *main_branches - - langfuse_logging_unit_tests: - filters: *main_branches - - litellm_assistants_api_testing: - filters: *main_branches - - litellm_router_testing: - filters: *main_branches - - litellm_router_unit_testing: - filters: *main_branches - - auth_ui_unit_tests: - filters: *main_branches - - build_docker_database_image: - filters: *main_branches - - e2e_ui_testing: - filters: *main_branches - - e2e_ui_testing_server_root_path: - filters: *main_branches + - base_sdk_install + - local_testing_part1 + - local_testing_part2 + - langfuse_logging_unit_tests + - litellm_assistants_api_testing + - litellm_router_testing + - litellm_router_unit_testing + - auth_ui_unit_tests + - build_docker_database_image + - e2e_ui_testing + - e2e_ui_testing_server_root_path - build_and_test: requires: - build_docker_database_image - filters: *main_branches - e2e_openai_endpoints: requires: - build_docker_database_image - filters: *main_branches - proxy_logging_guardrails_model_info_tests: requires: - build_docker_database_image - filters: *main_branches - proxy_spend_accuracy_tests: requires: - build_docker_database_image - filters: *main_branches - proxy_multi_instance_tests: requires: - build_docker_database_image - filters: *main_branches - proxy_store_model_in_db_tests: requires: - build_docker_database_image - filters: *main_branches - - proxy_build_from_pip_tests: - filters: *main_branches + - proxy_build_from_pip_tests - proxy_pass_through_endpoint_tests: requires: - build_docker_database_image - filters: *main_branches - proxy_e2e_anthropic_messages_tests: requires: - build_docker_database_image - filters: *main_branches - - llm_translation_testing: - filters: *main_branches - - realtime_translation_testing: - filters: *main_branches - - agent_testing: - filters: *main_branches - - guardrails_testing: - filters: *main_branches - - google_generate_content_endpoint_testing: - filters: *main_branches - - llm_responses_api_testing: - filters: *main_branches - - ocr_testing: - filters: *main_branches - - search_testing: - filters: *main_branches - - batches_testing: - filters: *main_branches - - litellm_utils_testing: - filters: *main_branches - - pass_through_unit_testing: - filters: *main_branches - - image_gen_testing: - filters: *main_branches - - logging_testing: - filters: *main_branches - - audio_testing: - filters: *main_branches - - redis_caching_unit_tests: - filters: *main_branches + - llm_translation_testing + - realtime_translation_testing + - agent_testing + - guardrails_testing + - google_generate_content_endpoint_testing + - llm_responses_api_testing + - ocr_testing + - search_testing + - batches_testing + - litellm_utils_testing + - pass_through_unit_testing + - image_gen_testing + - logging_testing + - audio_testing + - redis_caching_unit_tests - upload-coverage: requires: - realtime_translation_testing @@ -3469,18 +3420,12 @@ workflows: - db_migration_disable_update_check: requires: - build_docker_database_image - filters: *main_branches - - installing_litellm_on_python: - filters: *main_branches - - installing_litellm_on_python_3_13: - filters: *main_branches - - installing_litellm_on_python_v2_migration_resolver: - filters: *main_branches + - installing_litellm_on_python + - installing_litellm_on_python_3_13 + - installing_litellm_on_python_v2_migration_resolver - helm_chart_testing: requires: - build_docker_database_image - filters: *main_branches - test_bad_database_url: requires: - build_docker_database_image - filters: *main_branches diff --git a/.circleci/scripts/run_integration.sh b/.circleci/scripts/run_integration.sh index b617a79946c..984419717a3 100644 --- a/.circleci/scripts/run_integration.sh +++ b/.circleci/scripts/run_integration.sh @@ -26,6 +26,7 @@ guard_created=false guard_installed=false guard6_created=false guard6_installed=false +egress_cgroup=litellm-integration cleanup() { original_status=$? trap - EXIT INT TERM @@ -47,14 +48,14 @@ cleanup() { fi done if [ "$guard_installed" = true ]; then - sudo iptables -D OUTPUT -m owner --uid-owner "$(id -u)" -j integration_only || original_status=1 + sudo iptables -D OUTPUT -m cgroup --path "$egress_cgroup" -j integration_only || original_status=1 fi if [ "$guard_created" = true ]; then sudo iptables -F integration_only || original_status=1 sudo iptables -X integration_only || original_status=1 fi if [ "$guard6_installed" = true ]; then - sudo ip6tables -D OUTPUT -m owner --uid-owner "$(id -u)" -j integration_only || original_status=1 + sudo ip6tables -D OUTPUT -m cgroup --path "$egress_cgroup" -j integration_only || original_status=1 fi if [ "$guard6_created" = true ]; then sudo ip6tables -F integration_only || original_status=1 @@ -100,6 +101,8 @@ if [ "$mode" = parity ]; then export INTEGRATION_ROUTING=capture fi +sudo mkdir -p "/sys/fs/cgroup/$egress_cgroup" +echo "$$" | sudo tee "/sys/fs/cgroup/$egress_cgroup/cgroup.procs" > /dev/null sudo iptables -N integration_only guard_created=true sudo iptables -A integration_only -o lo -j ACCEPT @@ -109,13 +112,13 @@ for service in postgres-db redis-cache; do sudo iptables -A integration_only -d "$address" -j ACCEPT done sudo iptables -A integration_only -j REJECT -sudo iptables -I OUTPUT 1 -m owner --uid-owner "$(id -u)" -j integration_only +sudo iptables -I OUTPUT 1 -m cgroup --path "$egress_cgroup" -j integration_only guard_installed=true sudo ip6tables -N integration_only guard6_created=true sudo ip6tables -A integration_only -o lo -j ACCEPT sudo ip6tables -A integration_only -j REJECT -sudo ip6tables -I OUTPUT 1 -m owner --uid-owner "$(id -u)" -j integration_only +sudo ip6tables -I OUTPUT 1 -m cgroup --path "$egress_cgroup" -j integration_only guard6_installed=true if curl --noproxy '*' --connect-timeout 2 -s http://198.51.100.1 >/dev/null 2>&1; then diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 1fe0c602036..db46114715d 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -65,7 +65,7 @@ After: the same request comes back with real token counts, so the dashboard show **Please complete all items before asking a LiteLLM maintainer to review your PR** - [ ] I have added meaningful tests -- [ ] The handful of test files covering my change pass locally, e.g. `uv run pytest tests/test_litellm/.py -v`. Leave the suites (`make test-unit-*`, `make test-unit`) to CI: it finishes in ~15 minutes where a laptop takes an hour or more +- [ ] The handful of test files covering my change pass locally, e.g. `uv run pytest tests/unit/.py -v`. Leave the suites (`make test-unit-*`, `make test-unit`) to CI: it finishes in ~15 minutes where a laptop takes an hour or more - [ ] My PR passes all required CI/CD checks (e.g., lint, schema.d.ts sync check, etc.) - [ ] My PR's scope is as isolated as possible; it only solves 1 specific problem - [ ] I have received a Greptile **Confidence Score of at least 4/5** before requesting a maintainer review (Greptile reviews automatically once the PR is opened; only comment `@greptileai` to re-request a review after pushing changes) diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index 808bb2afd08..2d399cca3a4 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -14,7 +14,6 @@ on: - "litellm/ocr/**" - "litellm/llms/base_llm/ocr/**" - "litellm/llms/custom_httpx/llm_http_handler.py" - - "tests/test_litellm/ocr/**" - "tests/test_litellm/conftest.py" - "Makefile" - ".cargo/**" @@ -42,7 +41,6 @@ on: - "litellm/ocr/**" - "litellm/llms/base_llm/ocr/**" - "litellm/llms/custom_httpx/llm_http_handler.py" - - "tests/test_litellm/ocr/**" - "tests/test_litellm/conftest.py" - "Makefile" - ".cargo/**" diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index d75213d37ea..f55e186e3b2 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -61,7 +61,7 @@ jobs: - shard: core-utils artifact-name: core-utils - test-path: "tests/test_litellm/litellm_core_utils" + test-path: "" unit-flag: core-utils workers: 2 reruns: 1 @@ -88,7 +88,7 @@ jobs: - shard: Vertex AI artifact-name: llm-vertex-ai - test-path: "tests/test_litellm/llms/vertex_ai" + test-path: "" unit-flag: llm-vertex-ai workers: 1 reruns: 2 @@ -97,7 +97,7 @@ jobs: - shard: All Other Providers artifact-name: llm-other-providers - test-path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai" + test-path: "" unit-flag: llm-other-providers workers: 2 reruns: 2 @@ -107,9 +107,6 @@ jobs: - shard: misc artifact-name: misc test-path: >- - tests/test_litellm/interactions - tests/test_litellm/ocr - tests/test_litellm/passthrough tests/test_litellm/test_*.py unit-flag: misc workers: 2 diff --git a/AGENTS.md b/AGENTS.md index 69e034fbdea..a2dcd24bdd1 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -27,7 +27,7 @@ Never test structure of code only function of it A test must only fail when litellm code changes. Never pin facts we don't own (a vendor's price, a third party's field, an upstream default, today's date) as literals or as "X must be absent"; assert the invariant our code guarantees instead, e.g. two rows agree, a value is within range, a field is derived from another. If an outside fact is truly load-bearing, cite its source and date next to the assertion so a reader can tell stale from broken -`tests/test_litellm/` mirrors `litellm/` in a parallel path (see `tests/test_litellm/readme.md`). Name tests `test_.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_.py` if you're the first test there). One focused regression test beats many shallow ones +`tests/unit/` mirrors `litellm/` in a parallel path (see `tests/unit/AGENTS.md`). Name tests `test_.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_.py` if you're the first test there). One focused regression test beats many shallow ones End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `AGENTS.md` diff --git a/ARCHITECTURE.md b/ARCHITECTURE.md index f418752d990..c9e046748e8 100644 --- a/ARCHITECTURE.md +++ b/ARCHITECTURE.md @@ -255,7 +255,7 @@ Conventions to follow when touching this layer: | Column vs. field names | Where a model field differs from its DB column (for example `org_id` maps to the `organization_id` column), the repository translates in both directions rather than relying on Pydantic to guess. | | Array mutations | Adds use Prisma's atomic `push` (`add_member`, `add_admin`, `add_models`) to avoid read-modify-write races. Removals fall back to read-modify-write because Prisma has no atomic array remove. | -To add a new entity, define the model under `litellm/models/`, re-export it from `proxy/_types.py` if existing code imports it from there, and add a repository under `litellm/repositories/` (subclass `BaseRepository` for plain CRUD, or add bespoke methods when the entity needs encryption, archiving, or atomic array updates). Mirror the tests in `tests/test_litellm/repositories/`. +To add a new entity, define the model under `litellm/models/`, re-export it from `proxy/_types.py` if existing code imports it from there, and add a repository under `litellm/repositories/` (subclass `BaseRepository` for plain CRUD, or add bespoke methods when the entity needs encryption, archiving, or atomic array updates). Mirror the tests in `tests/unit/repositories/`. --- diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 082b7a8fb3e..a5ad6e97f3d 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -14,7 +14,7 @@ Here are the core requirements for any PR submitted to LiteLLM: - [ ] **Add testing** - Adding at least 1 test is a hard requirement - [see details](#adding-testing) - [ ] **Ensure your PR passes all checks**: - [ ] [Linting / Formatting](#running-linting-and-formatting-checks) - `make lint` - - [ ] [The tests covering your change](#running-unit-tests) pass, e.g. `uv run pytest tests/test_litellm/.py -v`. CI runs the full unit test matrix, so you don't need to run the whole suite locally + - [ ] [The tests covering your change](#running-unit-tests) pass, e.g. `uv run pytest tests/unit/.py -v`. CI runs the full unit test matrix, so you don't need to run the whole suite locally #### UI PRs @@ -72,7 +72,7 @@ make format make lint # Run the tests covering your change (CI runs the full suite) -uv run pytest tests/test_litellm/.py -v +uv run pytest tests/unit/.py -v # Commit your changes (must follow Conventional Commits — see above) git add . @@ -88,7 +88,7 @@ git push origin feature/your-feature ### Where to Add Tests -Add your tests to the [`tests/test_litellm/` directory](https://github.com/BerriAI/litellm/tree/main/tests/test_litellm). +Add your tests to the [`tests/unit/` directory](https://github.com/BerriAI/litellm/tree/main/tests/unit). - This directory mirrors the structure of the `litellm/` directory - **Only add mocked tests** - no real LLM API calls in this directory @@ -96,10 +96,10 @@ Add your tests to the [`tests/test_litellm/` directory](https://github.com/Berri ### File Naming Convention -The `tests/test_litellm/` directory follows the same structure as `litellm/`: +The `tests/unit/` directory follows the same structure as `litellm/`: - `litellm/proxy/caching_routes.py` → `tests/test_litellm/proxy/test_caching_routes.py` -- `litellm/utils.py` → `tests/test_litellm/test_utils.py` +- `litellm/utils.py` → `tests/unit/test_utils.py` ### Example Test @@ -125,10 +125,10 @@ def test_your_feature(): Run the tests covering your change: ```bash -uv run pytest tests/test_litellm/test_your_file.py -v +uv run pytest tests/unit/test_your_file.py -v ``` -`tests/test_litellm` holds thousands of tests, so running all of it locally takes a long time. CI runs it as a parallel matrix (`make test-unit-llms`, `make test-unit-proxy-core`, and the other `test-unit-*` targets) on beefier boxes, so if, for whatever reason, you must run the whole suite, it's better to rely on CI to do that. +`tests/unit` holds thousands of tests, so running all of it locally takes a long time. CI runs it as a parallel matrix (`make test-unit-llms`, `make test-unit-proxy-core`, and the other `test-unit-*` targets) on beefier boxes, so if, for whatever reason, you must run the whole suite, it's better to rely on CI to do that. If you're running broader test suites, proxy tests, or anything that touches PostgreSQL-backed fixtures/plugins, install the full local test environment first: diff --git a/Makefile b/Makefile index 311a7daef92..79c18f6fe82 100644 --- a/Makefile +++ b/Makefile @@ -42,7 +42,7 @@ help: @echo " make check-circular-imports - Check for circular imports" @echo " make check-import-safety - Check import safety" @echo " make test - Run all tests" - @echo " make test-unit - Run unit tests (tests/test_litellm)" + @echo " make test-unit - Run unit tests (tests/unit and tests/test_litellm)" @echo " make test-unit-llms - Run LLM provider tests (~225 files)" @echo " make test-unit-proxy-guardrails - Run proxy guardrails+mgmt tests (~51 files)" @echo " make test-unit-proxy-core - Run proxy auth+client+db+hooks tests (~52 files)" @@ -310,7 +310,7 @@ test: install-test-deps $(UV_RUN) pytest tests/ test-unit: install-test-deps - $(UV_RUN) pytest tests/test_litellm -x -vv -n 4 + $(UV_RUN) pytest tests/unit tests/test_litellm -x -vv -n 4 # Matrix test targets (matching CI workflow groups) test-unit-llms: install-test-deps @@ -332,7 +332,7 @@ test-unit-core-utils: install-test-deps $(UV_RUN) pytest tests/unit/litellm_core_utils --tb=short -vv -n 2 --durations=20 test-unit-other: install-test-deps - $(UV_RUN) pytest tests/unit/caching tests/unit/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/unit/router_strategy tests/unit/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20 + $(UV_RUN) pytest tests/unit/caching tests/unit/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/unit/router_strategy tests/unit/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20 test-unit-root: install-test-deps $(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20 diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py index 4be09670e92..6e33d9f1bf3 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py @@ -32,6 +32,7 @@ from litellm.integrations.email_templates.key_rotated_email import ( from litellm.integrations.email_templates.templates import ( MAX_BUDGET_ALERT_EMAIL_TEMPLATE, SOFT_BUDGET_ALERT_EMAIL_TEMPLATE, + TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE, TEAM_SOFT_BUDGET_ALERT_EMAIL_TEMPLATE, ) from litellm.integrations.email_templates.user_invitation_email import ( @@ -48,6 +49,12 @@ from litellm.secret_managers.main import get_secret_bool from litellm.types.integrations.slack_alerting import LITELLM_LOGO_URL +def _max_budget_alert_id(user_info: CallInfo) -> str: + if user_info.event_group == Litellm_EntityType.TEAM_MEMBER: + return f"team_member:{user_info.user_id}:{user_info.team_id}" + return user_info.token or user_info.user_id or "default_id" + + def _parse_email_list(raw) -> List[str]: """Parse emails from a list or comma-separated string.""" if isinstance(raw, list): @@ -373,17 +380,31 @@ class BaseEmailLogger(CustomLogger): greeting = html.escape( event.user_email or event.key_alias or event.token or "" ) - email_html_content = MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format( - email_logo_url=email_params.logo_url, - recipient_email=greeting, - percentage=percentage, - spend=spend_str, - max_budget=max_budget_str, - alert_threshold=alert_threshold_str, - base_url=email_params.base_url, - email_support_contact=email_params.support_contact, - email_footer=email_params.signature, - ) + if event.event_group == Litellm_EntityType.TEAM_MEMBER: + email_html_content = TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format( + email_logo_url=email_params.logo_url, + member=html.escape(event.user_email or event.user_id or ""), + team_alias=html.escape(event.team_alias or event.team_id or ""), + percentage=percentage, + spend=spend_str, + max_budget=max_budget_str, + alert_threshold=alert_threshold_str, + base_url=email_params.base_url, + email_support_contact=email_params.support_contact, + email_footer=email_params.signature, + ) + else: + email_html_content = MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format( + email_logo_url=email_params.logo_url, + recipient_email=greeting, + percentage=percentage, + spend=spend_str, + max_budget=max_budget_str, + alert_threshold=alert_threshold_str, + base_url=email_params.base_url, + email_support_contact=email_params.support_contact, + email_footer=email_params.signature, + ) await self.send_email( from_email=self.DEFAULT_LITELLM_EMAIL, to_email=recipient_emails, @@ -607,7 +628,7 @@ class BaseEmailLogger(CustomLogger): if user_info.spend < threshold_amount: continue - _id = user_info.token or user_info.user_id or "default_id" + _id = _max_budget_alert_id(user_info) _cache_key = ( f"email_budget_alerts:max_budget_alert:{threshold_pct}:{_id}" ) @@ -618,7 +639,7 @@ class BaseEmailLogger(CustomLogger): emails.append(user_info.user_email) if not emails: verbose_proxy_logger.warning( - "No recipients for %d%% threshold on key %s, skipping alert", + "No recipients for %d%% threshold on %s, skipping alert", threshold_pct, _id, ) @@ -633,7 +654,11 @@ class BaseEmailLogger(CustomLogger): if send_count is not None and send_count > 1: continue - event_message = f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached" + event_message = ( + f"Team Member Budget Alert - {threshold_pct}% of Team Member Budget Reached" + if user_info.event_group == Litellm_EntityType.TEAM_MEMBER + else f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached" + ) webhook_event = WebhookEvent( event="max_budget_alert", event_message=event_message, diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index 70fcc367905..bc6a2552e4c 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -9,10 +9,14 @@ - A test for another crate's item belongs in that crate, not in a downstream one - Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests//mod.rs` or `tests//support.rs` so it is not picked up as a test crate of its own +## Test fixtures and cases + +Use [`#[rstest]`](https://docs.rs/rstest/latest/rstest/attr.rstest.html) for new and updated tests and [`#[fixture]`](https://docs.rs/rstest/latest/rstest/attr.fixture.html) for reusable setup, injected through typed test arguments. Express input variations as named `#[case::name(...)]` cases instead of loops or duplicated tests so each failure identifies its case. Keep behavior assertions in the test body and fixtures focused on setup. Use the workspace `rstest` dependency + ## Error definitions - A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs` -- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each +- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each. A failure mode is something a caller handles differently (phase, status code, retry, a message Python parity pins exactly); failures no caller tells apart share one variant and differ only in its message - Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string - Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return - Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 3677d1d654f..d67623feffd 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -199,7 +199,7 @@ checksum = "ae36dc4177970ef04fde5178d3e2429882def40e57a451f919c098f72baa6cec" dependencies = [ "proc-macro2", "quote", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -710,14 +710,20 @@ dependencies = [ "http 1.4.2", "http-body 1.1.0", "http-body-util", + "hyper 1.10.1", + "hyper-util", "itoa", "matchit", "memchr", "mime", + "multer", "percent-encoding", "pin-project-lite", "serde_core", + "serde_json", + "serde_path_to_error", "sync_wrapper", + "tokio", "tower", "tower-layer", "tower-service", @@ -1053,18 +1059,18 @@ dependencies = [ [[package]] name = "clap" -version = "4.6.6" +version = "4.6.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "473c7e07f409a8d772161724aa8db6a765a2532a70f9667eeb7b49d3d02fbdca" +checksum = "aa8876b300ab35ba921adea3dfd70157a46249b33f95c9084ae5709785478946" dependencies = [ "clap_builder", ] [[package]] name = "clap_builder" -version = "4.6.6" +version = "4.6.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7b48fea5a88e9ae728a2dcbedbfc0e730f7d60da42e1cb049a83c9fb8b789889" +checksum = "ec0797fb7aeb1406c84efac526901f7ec3ead2124f946b494e72879d4b54704d" dependencies = [ "anstyle", "clap_lex", @@ -1180,7 +1186,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e75b2483e97a5a7da73ac68a05b629f9c53cff58d8ed1c77866079e18b00dba5" dependencies = [ "digest 0.10.7", - "spin", + "spin 0.10.1", ] [[package]] @@ -1581,6 +1587,15 @@ dependencies = [ "serde", ] +[[package]] +name = "encoding_rs" +version = "0.8.35" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3" +dependencies = [ + "cfg-if", +] + [[package]] name = "equivalent" version = "1.0.2" @@ -2816,6 +2831,10 @@ version = "0.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" +[[package]] +name = "litellm" +version = "0.0.1" + [[package]] name = "litellm-auth" version = "0.1.0" @@ -2840,6 +2859,7 @@ dependencies = [ "litellm-http", "moka", "reqwest 0.12.28", + "rstest", "serde_json", "sha2 0.10.9", "thiserror 2.0.19", @@ -2911,6 +2931,7 @@ dependencies = [ "litellm-cache", "litellm-cache-response", "litellm-cache-testing", + "litellm-http", "reqwest 0.12.28", "rstest", "serde_json", @@ -2944,6 +2965,7 @@ dependencies = [ "litellm-auth-types", "litellm-cache", "litellm-cache-testing", + "litellm-http", "percent-encoding", "reqwest 0.12.28", "rstest", @@ -2970,6 +2992,7 @@ dependencies = [ "futures-util", "litellm-cache", "litellm-cache-testing", + "litellm-http", "qdrant-client", "reqwest 0.12.28", "rstest", @@ -3043,6 +3066,7 @@ dependencies = [ "litellm-auth-aws", "litellm-cache", "litellm-cache-testing", + "litellm-http", "reqwest 0.12.28", "rstest", "serde_json", @@ -3087,6 +3111,18 @@ dependencies = [ "strum", ] +[[package]] +name = "litellm-config" +version = "0.1.0" +dependencies = [ + "litellm-auth-types", + "rstest", + "serde", + "serde_yaml_ng", + "tempfile", + "thiserror 2.0.19", +] + [[package]] name = "litellm-core" version = "0.1.0" @@ -3102,6 +3138,7 @@ dependencies = [ "litellm-http", "litellm-llms", "litellm-secrets", + "litellm-tracing", "litellm-types", "mime_guess", "moka", @@ -3137,6 +3174,7 @@ dependencies = [ "serde_json", "serde_path_to_error", "serde_with", + "strum", "thiserror 2.0.19", "url", ] @@ -3173,6 +3211,69 @@ dependencies = [ "tokio-util", ] +[[package]] +name = "litellm-gateway" +version = "0.1.0" +dependencies = [ + "axum", + "futures-util", + "http-body-util", + "litellm-config", + "litellm-core", + "litellm-gateway-auth", + "litellm-gateway-inference", + "litellm-http", + "litellm-llms", + "litellm-secrets", + "litellm-tracing", + "rstest", + "serde_json", + "tokio", + "tower", + "tracing", + "uuid", +] + +[[package]] +name = "litellm-gateway-auth" +version = "0.1.0" +dependencies = [ + "axum", + "futures-util", + "litellm-auth-types", + "litellm-config", + "litellm-secrets", + "rstest", + "sha2 0.10.9", + "subtle", + "thiserror 2.0.19", + "tokio", + "tower", +] + +[[package]] +name = "litellm-gateway-inference" +version = "0.1.0" +dependencies = [ + "axum", + "base64 0.22.1", + "bytes", + "futures-util", + "litellm-auth", + "litellm-core", + "litellm-http", + "litellm-llms", + "litellm-router", + "litellm-secrets", + "litellm-types", + "rstest", + "serde_json", + "thiserror 2.0.19", + "tokio", + "tower", + "wiremock", +] + [[package]] name = "litellm-host" version = "0.1.0" @@ -3207,11 +3308,13 @@ dependencies = [ "http 1.4.2", "hyper-util", "litellm-core-utils", + "rcgen", "reqwest 0.12.28", "rstest", "rustls 0.23.42", "serde", "serde_json", + "tempfile", "thiserror 2.0.19", "tokio", "veil", @@ -3236,6 +3339,7 @@ dependencies = [ "litellm-framing", "litellm-host", "litellm-http", + "litellm-python-compat", "litellm-secrets", "litellm-types", "reqwest 0.12.28", @@ -3257,6 +3361,7 @@ version = "0.1.0" dependencies = [ "indexmap 2.14.0", "jsonschema", + "litellm-types", "rstest", "schemars 1.2.2", "serde", @@ -3276,7 +3381,6 @@ dependencies = [ "futures-util", "litellm-auth", "litellm-auth-aws", - "litellm-auth-gcp", "litellm-cache", "litellm-cache-azure-blob", "litellm-cache-disk", @@ -3311,6 +3415,7 @@ dependencies = [ "serde_json", "serde_with", "sha2 0.10.9", + "strum", "thiserror 2.0.19", "tokio", "tokio-tungstenite", @@ -3334,6 +3439,15 @@ dependencies = [ "thiserror 2.0.19", ] +[[package]] +name = "litellm-router" +version = "0.1.0" +dependencies = [ + "litellm-config", + "litellm-core", + "rstest", +] + [[package]] name = "litellm-secrets" version = "0.1.0" @@ -3344,6 +3458,7 @@ dependencies = [ "google-cloud-auth", "google-cloud-kms-v1", "litellm-core-utils", + "litellm-http", "litellm-python-compat", "litellm-secrets-aws", "litellm-secrets-azure", @@ -3391,6 +3506,7 @@ dependencies = [ "litellm-auth-azure", "litellm-auth-types", "litellm-core-utils", + "litellm-http", "litellm-secrets-types", "percent-encoding", "reqwest 0.12.28", @@ -3410,6 +3526,7 @@ version = "0.1.0" dependencies = [ "base64 0.22.1", "litellm-core-utils", + "litellm-http", "litellm-secrets-types", "litellm-tracing", "moka", @@ -3438,6 +3555,7 @@ dependencies = [ "litellm-auth-gcp", "litellm-auth-types", "litellm-core-utils", + "litellm-http", "litellm-secrets-types", "moka", "percent-encoding", @@ -3479,6 +3597,7 @@ dependencies = [ "rstest", "serde", "serde_json", + "strum", "thiserror 2.0.19", "tokio", "veil", @@ -3562,6 +3681,7 @@ dependencies = [ name = "litellm-tracing" version = "0.1.0" dependencies = [ + "base64 0.22.1", "fancy-regex 0.19.2", "percent-encoding", "rstest", @@ -3576,8 +3696,10 @@ name = "litellm-types" version = "0.1.0" dependencies = [ "rstest", + "schemars 1.2.2", "serde", "serde_json", + "strum", ] [[package]] @@ -3745,6 +3867,23 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "multer" +version = "3.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "83e87776546dc87511aa5ee218730c92b666d7264ab6ed41f9d215af9cd5224b" +dependencies = [ + "bytes", + "encoding_rs", + "futures-util", + "http 1.4.2", + "httparse", + "memchr", + "mime", + "spin 0.9.9", + "version_check", +] + [[package]] name = "nom" version = "7.1.3" @@ -4652,7 +4791,7 @@ checksum = "92ecd8964f8453721699a1ed72037b0db49ce2f5a5138486ee89bed6f67cdf3a" dependencies = [ "proc-macro2", "quote", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -5131,7 +5270,7 @@ dependencies = [ "proc-macro2", "quote", "serde_derive_internals", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -5219,7 +5358,7 @@ checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -5230,7 +5369,7 @@ checksum = "f852137cce035d6a4df67ccce505ff6b3e9fd3a10e3e52b24dc71e650bb1a9bd" dependencies = [ "proc-macro2", "quote", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -5310,6 +5449,19 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "serde_yaml_ng" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b4db627b98b36d4203a7b458cf3573730f2bb591b28871d916dfa9efabfd41f" +dependencies = [ + "indexmap 2.14.0", + "itoa", + "ryu", + "serde", + "unsafe-libyaml", +] + [[package]] name = "sha1" version = "0.10.7" @@ -5439,6 +5591,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "spin" +version = "0.9.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e" + [[package]] name = "spin" version = "0.10.1" @@ -5538,9 +5696,9 @@ dependencies = [ [[package]] name = "syn" -version = "3.0.0" +version = "3.0.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f2fac314a64dc9a36e61a9eb4261a5e9bbfbc922b27e518af97bc32b926cf967" +checksum = "8593e8e72159ed2257d083c7a454a85cbf854f37a0966d8d483aff8c8a3ebcee" dependencies = [ "proc-macro2", "quote", @@ -5652,7 +5810,7 @@ checksum = "43cbfe0cf76104d42a574802844187e84a305e531ed54455f11fbde0f10541cd" dependencies = [ "proc-macro2", "quote", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -6231,6 +6389,12 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e" +[[package]] +name = "unsafe-libyaml" +version = "0.2.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "673aac59facbab8a9007c7f6108d11f63b603f7cabff99fabf650fea5c32b861" + [[package]] name = "untrusted" version = "0.9.0" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 022e8f13311..ed703396c22 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -9,9 +9,13 @@ license = "MIT" repository = "https://github.com/BerriAI/litellm" [workspace.dependencies] +litellm-config = { path = "crates/config" } +litellm-router = { path = "crates/router" } litellm-tracing = { path = "crates/tracing" } -tracing = "0.1" litellm-core = { path = "crates/core" } +litellm-gateway = { path = "crates/gateway" } +litellm-gateway-inference = { path = "crates/gateway-inference" } +litellm-gateway-auth = { path = "crates/gateway-auth" } litellm-coroutine = { path = "crates/coroutine" } litellm-host = { path = "crates/host" } litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" } @@ -48,7 +52,10 @@ litellm-token-counter-fast = { path = "crates/token-counter-fast" } litellm-token-counter-huggingface = { path = "crates/token-counter-huggingface" } litellm-token-counter-tiktoken = { path = "crates/token-counter-tiktoken" } litellm-host-python = { path = "crates/host-python" } +litellm-python-compat = { path = "crates/python-compat" } +tracing = "0.1" +axum = { version = "0.8.9", default-features = false, features = ["http1", "tokio", "multipart"] } bytes = "1" http = "1" google-cloud-auth = { version = "1.16.0", default-features = false } @@ -57,8 +64,8 @@ hyper-util = { version = "0.1.20", default-features = false, features = ["client proptest = "1.7.0" pyo3 = "0.29.2" pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] } -pythonize = "0.29.0" rand = "0.8" +schemars = "1" reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls", "http2", "stream"] } qdrant-client = { version = "1.19.0", default-features = false } uuid = { version = "1", features = ["v4"] } diff --git a/litellm-rust/clippy.toml b/litellm-rust/clippy.toml index f7e3293069b..0e2ff770d27 100644 --- a/litellm-rust/clippy.toml +++ b/litellm-rust/clippy.toml @@ -7,4 +7,16 @@ disallowed-methods = [ { path = "pyo3_async_runtimes::tokio::local_future_into_py", reason = "use litellm_host_python::run_async / run_async_value" }, { path = "pyo3_async_runtimes::tokio::run", reason = "use litellm_host_python::run_sync / run_sync_value" }, { path = "pyo3_async_runtimes::tokio::run_until_complete", reason = "use litellm_host_python::run_sync / run_sync_value" }, + { path = "reqwest::Client::new", reason = "take litellm_http::Client from HttpClientPool" }, + { path = "reqwest::Client::builder", reason = "HttpClientConfig owns client construction" }, + { path = "reqwest::ClientBuilder::danger_accept_invalid_certs", reason = "set HttpClientConfig::verify instead" }, + { path = "reqwest::ClientBuilder::identity", reason = "set HttpClientConfig::client_certificate instead" }, + { path = "reqwest::ClientBuilder::use_preconfigured_tls", reason = "HttpClientConfig owns the TLS configuration" }, +] + +# Every outbound client comes from litellm_http::HttpClientPool so it honors the host's TLS, +# proxy and timeout settings. Only crates/http builds one. +disallowed-types = [ + { path = "reqwest::Client", reason = "take litellm_http::Client from HttpClientPool; only crates/http builds one" }, + { path = "reqwest::ClientBuilder", reason = "HttpClientConfig owns client construction" }, ] diff --git a/litellm-rust/crates/auth-aws/Cargo.toml b/litellm-rust/crates/auth-aws/Cargo.toml index 1a35af48574..2a9a9e4768c 100644 --- a/litellm-rust/crates/auth-aws/Cargo.toml +++ b/litellm-rust/crates/auth-aws/Cargo.toml @@ -22,5 +22,7 @@ aws-types = "1.4.0" aws-smithy-runtime-api = "1.13.0" [dev-dependencies] +rstest.workspace = true +litellm-http = { workspace = true, features = ["test-support"] } reqwest.workspace = true tokio.workspace = true diff --git a/litellm-rust/crates/auth-aws/src/aws.rs b/litellm-rust/crates/auth-aws/src/aws.rs index cb9195ffeb6..69eb4265159 100644 --- a/litellm-rust/crates/auth-aws/src/aws.rs +++ b/litellm-rust/crates/auth-aws/src/aws.rs @@ -1,5 +1,4 @@ use std::collections::BTreeMap; -use std::sync::OnceLock; use std::time::Duration; use std::time::{SystemTime, UNIX_EPOCH}; @@ -26,8 +25,26 @@ use super::constants::{ const STATIC_CREDENTIALS_TTL: Duration = Duration::from_secs(3600 - 60); const AMBIENT_CREDENTIALS_TTL: Duration = Duration::from_secs(600); -static STATIC_CREDENTIALS_CACHE: OnceLock> = OnceLock::new(); -static AMBIENT_CREDENTIALS_CACHE: OnceLock> = OnceLock::new(); +#[derive(Clone)] +pub struct AwsAuthService { + static_credentials: Cache, + ambient_credentials: Cache, +} + +impl Default for AwsAuthService { + fn default() -> Self { + Self { + static_credentials: Cache::builder() + .max_capacity(200) + .time_to_live(STATIC_CREDENTIALS_TTL) + .build(), + ambient_credentials: Cache::builder() + .max_capacity(200) + .time_to_live(AMBIENT_CREDENTIALS_TTL) + .build(), + } + } +} fn credential_cache_ttl(flow: &AwsAuthFlow) -> Option { match flow { @@ -108,35 +125,19 @@ fn cache_key(config: &AwsAuthConfig, flow: &AwsAuthFlow) -> String { format!("{:x}", hasher.finalize()) } -fn static_credentials_cache() -> &'static Cache { - STATIC_CREDENTIALS_CACHE.get_or_init(|| { - Cache::builder() - .max_capacity(200) - .time_to_live(STATIC_CREDENTIALS_TTL) - .build() - }) -} +impl AwsAuthService { + fn get_cached_credentials(&self, key: &str) -> Option { + self.static_credentials + .get(key) + .or_else(|| self.ambient_credentials.get(key)) + } -fn ambient_credentials_cache() -> &'static Cache { - AMBIENT_CREDENTIALS_CACHE.get_or_init(|| { - Cache::builder() - .max_capacity(200) - .time_to_live(AMBIENT_CREDENTIALS_TTL) - .build() - }) -} - -fn get_cached_credentials(key: &str) -> Option { - static_credentials_cache() - .get(key) - .or_else(|| ambient_credentials_cache().get(key)) -} - -fn set_cached_credentials(key: String, credentials: Credentials, ttl: Duration) { - if ttl == STATIC_CREDENTIALS_TTL { - static_credentials_cache().insert(key, credentials); - } else { - ambient_credentials_cache().insert(key, credentials); + fn set_cached_credentials(&self, key: String, credentials: Credentials, ttl: Duration) { + if ttl == STATIC_CREDENTIALS_TTL { + self.static_credentials.insert(key, credentials); + } else { + self.ambient_credentials.insert(key, credentials); + } } } @@ -214,66 +215,157 @@ pub fn classify_auth( AwsAuthFlow::DefaultChain } -pub async fn resolve_credentials( - config: AwsAuthConfig, - env_lookup: &(dyn Fn(&str) -> Option + Sync), -) -> Result { - let resolved = config.clone().with_environment(env_lookup); - let flow = classify_auth(config, env_lookup); - match flow { - AwsAuthFlow::SessionToken { - access_key_id, - secret_access_key, - session_token, - } => Ok(Credentials::new( - access_key_id, - secret_access_key, - Some(session_token), - None, - "litellm-static-session", - )), - AwsAuthFlow::StaticKeys { - access_key_id, - secret_access_key, - region_name, - } => { - let flow = AwsAuthFlow::StaticKeys { - access_key_id: access_key_id.clone(), - secret_access_key: secret_access_key.clone(), - region_name, - }; - let key = cache_key(&resolved, &flow); - if let Some(credentials) = get_cached_credentials(&key) { - return Ok(credentials); - } - let credentials = Credentials::new( +impl AwsAuthService { + pub async fn resolve_credentials( + &self, + config: AwsAuthConfig, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + ) -> Result { + let resolved = config.clone().with_environment(env_lookup); + let flow = classify_auth(config, env_lookup); + match flow { + AwsAuthFlow::SessionToken { access_key_id, secret_access_key, + session_token, + } => Ok(Credentials::new( + access_key_id, + secret_access_key, + Some(session_token), None, - None, - "litellm-static", - ); - set_cached_credentials( - key, - credentials.clone(), - credential_cache_ttl(&flow).unwrap_or(STATIC_CREDENTIALS_TTL), - ); - Ok(credentials) - } - AwsAuthFlow::Profile { name } => { - let provider = aws_config::profile::ProfileFileCredentialsProvider::builder() - .profile_name(name) - .build(); - provider - .provide_credentials() - .await - .map_err(|error| Error::AwsProfile(error.to_string())) - } - AwsAuthFlow::AssumeRole { role, session_name } => { - if is_already_running_as_role(&role, &resolved).await? { - let ambient_flow = AwsAuthFlow::DefaultChain; - let key = cache_key(&resolved, &ambient_flow); - if let Some(credentials) = get_cached_credentials(&key) { + "litellm-static-session", + )), + AwsAuthFlow::StaticKeys { + access_key_id, + secret_access_key, + region_name, + } => { + let flow = AwsAuthFlow::StaticKeys { + access_key_id: access_key_id.clone(), + secret_access_key: secret_access_key.clone(), + region_name, + }; + let key = cache_key(&resolved, &flow); + if let Some(credentials) = self.get_cached_credentials(&key) { + return Ok(credentials); + } + let credentials = Credentials::new( + access_key_id, + secret_access_key, + None, + None, + "litellm-static", + ); + self.set_cached_credentials( + key, + credentials.clone(), + credential_cache_ttl(&flow).unwrap_or(STATIC_CREDENTIALS_TTL), + ); + Ok(credentials) + } + AwsAuthFlow::Profile { name } => { + let provider = aws_config::profile::ProfileFileCredentialsProvider::builder() + .profile_name(name) + .build(); + provider + .provide_credentials() + .await + .map_err(|error| Error::AwsProfile(error.to_string())) + } + AwsAuthFlow::AssumeRole { role, session_name } => { + if is_already_running_as_role(&role, &resolved).await? { + let ambient_flow = AwsAuthFlow::DefaultChain; + let key = cache_key(&resolved, &ambient_flow); + if let Some(credentials) = self.get_cached_credentials(&key) { + return Ok(credentials); + } + let provider = + aws_config::default_provider::credentials::DefaultCredentialsChain::builder() + .build() + .await; + let credentials = provider + .provide_credentials() + .await + .map_err(|error| Error::AwsDefaultChain(error.to_string()))?; + self.set_cached_credentials( + key, + credentials.clone(), + credential_cache_ttl(&ambient_flow).unwrap_or(AMBIENT_CREDENTIALS_TTL), + ); + return Ok(credentials); + } + let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest()); + if let Some(region) = resolved.region_name.clone() { + loader = loader.region(aws_types::region::Region::new(region)); + } + if let Some(endpoint) = resolved.sts_endpoint.clone() { + loader = loader.endpoint_url(endpoint); + } + if let (Some(access_key_id), Some(secret_access_key)) = + (resolved.access_key_id, resolved.secret_access_key) + { + loader = loader.credentials_provider(Credentials::new( + access_key_id, + secret_access_key, + resolved.session_token, + None, + "litellm-role-source", + )); + } + let sdk_config = loader.load().await; + let builder = aws_config::sts::AssumeRoleProvider::builder(role); + let builder = match session_name { + Some(name) => builder.session_name(name), + None => builder.session_name(default_session_name()), + }; + let builder = match resolved.external_id { + Some(id) => builder.external_id(id), + None => builder, + }; + let provider = builder.configure(&sdk_config).build().await; + provider + .provide_credentials() + .await + .map_err(|error| Error::AwsAssumeRole(error.to_string())) + } + AwsAuthFlow::WebIdentity { + token, + role, + session_name, + } => { + let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest()); + if let Some(region) = resolved.region_name { + loader = loader.region(aws_types::region::Region::new(region)); + } + if let Some(endpoint) = resolved.sts_endpoint { + loader = loader.endpoint_url(endpoint); + } + let sdk_config = loader.load().await; + let client = aws_sdk_sts::Client::new(&sdk_config); + let response = client + .assume_role_with_web_identity() + .role_arn(role) + .role_session_name(session_name) + .web_identity_token(token) + .send() + .await + .map_err(|error| Error::AwsWebIdentity(error.to_string()))?; + let credentials = response + .credentials() + .ok_or(Error::AwsMissingWebIdentityCredentials)?; + let expiration = SystemTime::try_from(*credentials.expiration()) + .map_err(|error| Error::AwsWebIdentityExpiration(error.to_string()))?; + Ok(Credentials::new( + credentials.access_key_id(), + credentials.secret_access_key(), + Some(credentials.session_token().to_string()), + Some(expiration), + "litellm-web-identity", + )) + } + AwsAuthFlow::DefaultChain => { + let key = cache_key(&resolved, &AwsAuthFlow::DefaultChain); + if let Some(credentials) = self.get_cached_credentials(&key) { return Ok(credentials); } let provider = @@ -284,101 +376,14 @@ pub async fn resolve_credentials( .provide_credentials() .await .map_err(|error| Error::AwsDefaultChain(error.to_string()))?; - set_cached_credentials( + self.set_cached_credentials( key, credentials.clone(), - credential_cache_ttl(&ambient_flow).unwrap_or(AMBIENT_CREDENTIALS_TTL), + credential_cache_ttl(&AwsAuthFlow::DefaultChain) + .unwrap_or(AMBIENT_CREDENTIALS_TTL), ); - return Ok(credentials); + Ok(credentials) } - let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest()); - if let Some(region) = resolved.region_name.clone() { - loader = loader.region(aws_types::region::Region::new(region)); - } - if let Some(endpoint) = resolved.sts_endpoint.clone() { - loader = loader.endpoint_url(endpoint); - } - if let (Some(access_key_id), Some(secret_access_key)) = - (resolved.access_key_id, resolved.secret_access_key) - { - loader = loader.credentials_provider(Credentials::new( - access_key_id, - secret_access_key, - resolved.session_token, - None, - "litellm-role-source", - )); - } - let sdk_config = loader.load().await; - let builder = aws_config::sts::AssumeRoleProvider::builder(role); - let builder = match session_name { - Some(name) => builder.session_name(name), - None => builder.session_name(default_session_name()), - }; - let builder = match resolved.external_id { - Some(id) => builder.external_id(id), - None => builder, - }; - let provider = builder.configure(&sdk_config).build().await; - provider - .provide_credentials() - .await - .map_err(|error| Error::AwsAssumeRole(error.to_string())) - } - AwsAuthFlow::WebIdentity { - token, - role, - session_name, - } => { - let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest()); - if let Some(region) = resolved.region_name { - loader = loader.region(aws_types::region::Region::new(region)); - } - if let Some(endpoint) = resolved.sts_endpoint { - loader = loader.endpoint_url(endpoint); - } - let sdk_config = loader.load().await; - let client = aws_sdk_sts::Client::new(&sdk_config); - let response = client - .assume_role_with_web_identity() - .role_arn(role) - .role_session_name(session_name) - .web_identity_token(token) - .send() - .await - .map_err(|error| Error::AwsWebIdentity(error.to_string()))?; - let credentials = response - .credentials() - .ok_or(Error::AwsMissingWebIdentityCredentials)?; - let expiration = SystemTime::try_from(*credentials.expiration()) - .map_err(|error| Error::AwsWebIdentityExpiration(error.to_string()))?; - Ok(Credentials::new( - credentials.access_key_id(), - credentials.secret_access_key(), - Some(credentials.session_token().to_string()), - Some(expiration), - "litellm-web-identity", - )) - } - AwsAuthFlow::DefaultChain => { - let key = cache_key(&resolved, &AwsAuthFlow::DefaultChain); - if let Some(credentials) = get_cached_credentials(&key) { - return Ok(credentials); - } - let provider = - aws_config::default_provider::credentials::DefaultCredentialsChain::builder() - .build() - .await; - let credentials = provider - .provide_credentials() - .await - .map_err(|error| Error::AwsDefaultChain(error.to_string()))?; - set_cached_credentials( - key, - credentials.clone(), - credential_cache_ttl(&AwsAuthFlow::DefaultChain).unwrap_or(AMBIENT_CREDENTIALS_TTL), - ); - Ok(credentials) } } } @@ -585,6 +590,37 @@ pub fn aws_auth_config( } } +/// Where the credentials that sign a request come from, decided when the request is +/// prepared and resolved when it is sent. +#[derive(Clone, Debug, PartialEq)] +pub enum AwsCredentialSource { + HostSupplied(Credentials), + Chain(AwsAuthConfig), +} + +impl AwsCredentialSource { + pub fn from_params( + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Self { + match host_supplied_credentials(optional_params) { + Some(credentials) => Self::HostSupplied(credentials), + None => Self::Chain(aws_auth_config(optional_params, env_lookup)), + } + } + + pub async fn resolve( + self, + auth: &AwsAuthService, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + ) -> Result { + match self { + Self::HostSupplied(credentials) => Ok(credentials), + Self::Chain(config) => auth.resolve_credentials(config, env_lookup).await, + } + } +} + /// Credentials a host resolved through its own chain and handed down verbatim. /// /// A host with its own resolution (LiteLLM's Python `BaseAWSLLM`, which reads @@ -747,17 +783,18 @@ mod tests { #[tokio::test] async fn static_credentials_do_not_use_network() { - let credentials = resolve_credentials( - AwsAuthConfig { - access_key_id: Some("ak".into()), - secret_access_key: Some("sk".into()), - region_name: Some("us-east-1".into()), - ..Default::default() - }, - &no_env, - ) - .await - .expect("static credentials"); + let credentials = AwsAuthService::default() + .resolve_credentials( + AwsAuthConfig { + access_key_id: Some("ak".into()), + secret_access_key: Some("sk".into()), + region_name: Some("us-east-1".into()), + ..Default::default() + }, + &no_env, + ) + .await + .expect("static credentials"); assert_eq!(credentials.access_key_id(), "ak"); assert_eq!(credentials.session_token(), None); } @@ -807,17 +844,67 @@ mod tests { ); } - #[test] + #[rstest::rstest] fn cache_round_trip_preserves_credentials() { + let auth = AwsAuthService::default(); let key = format!("cache-test-{}", std::process::id()); let credentials = Credentials::new("cache-ak", "cache-sk", None, None, "test"); - set_cached_credentials(key.clone(), credentials.clone(), STATIC_CREDENTIALS_TTL); + auth.set_cached_credentials(key.clone(), credentials.clone(), STATIC_CREDENTIALS_TTL); assert_eq!( - get_cached_credentials(&key).map(|value| value.access_key_id().to_string()), + auth.get_cached_credentials(&key) + .map(|value| value.access_key_id().to_string()), Some("cache-ak".to_string()) ); } + #[rstest::rstest] + #[tokio::test] + async fn cloned_services_reuse_credentials_but_independent_services_do_not() { + let auth = AwsAuthService::default(); + let config = AwsAuthConfig { + access_key_id: Some("configured-key".into()), + secret_access_key: Some("configured-secret".into()), + region_name: Some("us-east-1".into()), + ..AwsAuthConfig::default() + }; + let flow = classify_auth(config.clone(), &no_env); + let cached = Credentials::new("cached-key", "cached-secret", None, None, "test"); + auth.set_cached_credentials( + cache_key(&config, &flow), + cached.clone(), + STATIC_CREDENTIALS_TTL, + ); + + let reused = auth + .clone() + .resolve_credentials(config.clone(), &no_env) + .await + .unwrap(); + let independent = AwsAuthService::default() + .resolve_credentials(config.clone(), &no_env) + .await + .unwrap(); + let different = AwsAuthConfig { + access_key_id: Some("different-key".into()), + ..config.clone() + }; + let other_identity = auth + .resolve_credentials(different.clone(), &no_env) + .await + .unwrap(); + + assert_eq!(reused.access_key_id(), cached.access_key_id()); + assert_eq!(reused.secret_access_key(), cached.secret_access_key()); + assert_eq!( + Some(independent.access_key_id()), + config.access_key_id.as_deref() + ); + assert_eq!( + Some(other_identity.access_key_id()), + different.access_key_id.as_deref() + ); + } + #[test] fn same_role_comparison_matches_partition_account_and_role() { assert!(same_role_arns( @@ -952,17 +1039,18 @@ mod tests { let body = br#"{"anthropic_version":"bedrock-2023-05-31","max_tokens":1,"messages":[{"role":"user","content":[{"type":"text","text":"ping"}]}]}"#.to_vec(); let headers = BTreeMap::from([("Content-Type".to_string(), "application/json".to_string())]); - let credentials = resolve_credentials( - AwsAuthConfig { - access_key_id: Some(access_key_id), - secret_access_key: Some(secret_access_key), - region_name: Some("us-west-2".to_string()), - ..Default::default() - }, - &no_env, - ) - .await?; - let client = reqwest::Client::new(); + let credentials = AwsAuthService::default() + .resolve_credentials( + AwsAuthConfig { + access_key_id: Some(access_key_id), + secret_access_key: Some(secret_access_key), + region_name: Some("us-west-2".to_string()), + ..Default::default() + }, + &no_env, + ) + .await?; + let client = litellm_http::Client::plain_for_test(); let mut failures = Vec::new(); for region in ["us-west-2", "us-east-1"] { diff --git a/litellm-rust/crates/auth-aws/src/signer.rs b/litellm-rust/crates/auth-aws/src/signer.rs index 49a3910c1d5..3868fdd939b 100644 --- a/litellm-rust/crates/auth-aws/src/signer.rs +++ b/litellm-rust/crates/auth-aws/src/signer.rs @@ -1,13 +1,11 @@ use std::{collections::BTreeMap, time::SystemTime}; +use crate::{ + AwsAuthService, AwsCredentialSource, Error, aws_signature_headers, is_sigv4_computed_header, + sign_post, +}; use aws_credential_types::Credentials; use litellm_http::outbound::{RequestSigner, UnsignedRequest}; -use serde_json::{Map, Value}; - -use crate::{ - Error, aws_auth_config, aws_signature_headers, host_supplied_credentials, - is_sigv4_computed_header, resolve_credentials, sign_post, -}; #[derive(Clone, Debug)] pub struct SigV4Signer { @@ -32,19 +30,17 @@ impl SigV4Signer { } pub async fn resolve( + auth: &AwsAuthService, region: String, service: &'static str, - optional_params: &Map, + credentials: AwsCredentialSource, env_lookup: &(dyn Fn(&str) -> Option + Sync), ) -> Result { - let credentials = match host_supplied_credentials(optional_params) { - Some(credentials) => credentials, - None => { - resolve_credentials(aws_auth_config(optional_params, env_lookup), env_lookup) - .await? - } - }; - Ok(Self::new(region, service, credentials)) + Ok(Self::new( + region, + service, + credentials.resolve(auth, env_lookup).await?, + )) } } @@ -80,7 +76,7 @@ mod tests { use std::time::{Duration, UNIX_EPOCH}; use litellm_http::outbound::OutboundRequest; - use serde_json::json; + use serde_json::{Value, json}; use super::*; diff --git a/litellm-rust/crates/auth-gcp/src/lib.rs b/litellm-rust/crates/auth-gcp/src/lib.rs index 682f1af5fe1..4374dff95aa 100644 --- a/litellm-rust/crates/auth-gcp/src/lib.rs +++ b/litellm-rust/crates/auth-gcp/src/lib.rs @@ -131,7 +131,7 @@ impl Default for VertexAuth { } impl VertexAuth { - fn new(loader: Arc) -> Self { + pub fn new(loader: Arc) -> Self { Self { providers: Cache::builder().max_capacity(64).build(), loader, @@ -220,16 +220,16 @@ impl VertexAuth { } } -trait VertexTokenSource: Send + Sync { +pub trait VertexTokenSource: Send + Sync { fn project_id(&self) -> VertexAuthFuture<'_, String>; fn token(&self) -> VertexAuthFuture<'_, String>; } -trait VertexProviderLoader: Send + Sync { +pub trait VertexProviderLoader: Send + Sync { fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc>; } -type VertexAuthFuture<'a, T> = Pin> + Send + 'a>>; +pub type VertexAuthFuture<'a, T> = Pin> + Send + 'a>>; struct GcpTokenSource(Arc); @@ -305,7 +305,7 @@ fn validate_request_credentials(configured: &str) -> Result<&str, Error> { } #[derive(Clone, Debug)] -enum CredentialSource { +pub enum CredentialSource { Inline(SecretValue), Trusted(SecretValue), ApplicationCredentials(String), diff --git a/litellm-rust/crates/auth-types/src/http.rs b/litellm-rust/crates/auth-types/src/http.rs index 0cb5839f965..1519769521f 100644 --- a/litellm-rust/crates/auth-types/src/http.rs +++ b/litellm-rust/crates/auth-types/src/http.rs @@ -7,7 +7,7 @@ pub enum CredentialPlacement { } impl CredentialPlacement { - pub fn header_name(self) -> &'static str { + pub const fn header_name(self) -> &'static str { match self { Self::Bearer => "Authorization", Self::Header(name) => name, @@ -40,21 +40,6 @@ pub fn apply_credential( ) } -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum RequestAuth { - Header { - name: &'static str, - value: String, - }, - Bearer { - token: String, - }, - AwsSigV4 { - region: String, - service: &'static str, - }, -} - #[cfg(test)] mod tests { use super::{CredentialPlacement, apply_credential}; diff --git a/litellm-rust/crates/auth-types/src/lib.rs b/litellm-rust/crates/auth-types/src/lib.rs index 9d399249c05..ebab5b84d88 100644 --- a/litellm-rust/crates/auth-types/src/lib.rs +++ b/litellm-rust/crates/auth-types/src/lib.rs @@ -51,7 +51,7 @@ pub use credential::{ CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle, }; pub use error::Error; -pub use http::{CredentialPlacement, RequestAuth}; +pub use http::CredentialPlacement; pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy}; pub use secret::SecretValue; pub use token::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; diff --git a/litellm-rust/crates/auth/src/lib.rs b/litellm-rust/crates/auth/src/lib.rs index 622a5b2d58b..d23bccccc9f 100644 --- a/litellm-rust/crates/auth/src/lib.rs +++ b/litellm-rust/crates/auth/src/lib.rs @@ -2,6 +2,9 @@ pub use litellm_auth_types::*; +mod services; +pub use services::AuthServices; + #[cfg(feature = "aws")] pub use litellm_auth_aws as aws; #[cfg(feature = "azure")] diff --git a/litellm-rust/crates/auth/src/services.rs b/litellm-rust/crates/auth/src/services.rs new file mode 100644 index 00000000000..4c88c9a89a2 --- /dev/null +++ b/litellm-rust/crates/auth/src/services.rs @@ -0,0 +1,9 @@ +#[derive(Default)] +pub struct AuthServices { + #[cfg(feature = "aws")] + pub aws: litellm_auth_aws::AwsAuthService, + #[cfg(feature = "azure")] + pub azure: litellm_auth_azure::AzureAuthService, + #[cfg(feature = "gcp")] + pub gcp: litellm_auth_gcp::VertexAuth, +} diff --git a/litellm-rust/crates/cache-azure-blob/Cargo.toml b/litellm-rust/crates/cache-azure-blob/Cargo.toml index baa1b0f5482..5bdfa16ef53 100644 --- a/litellm-rust/crates/cache-azure-blob/Cargo.toml +++ b/litellm-rust/crates/cache-azure-blob/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true litellm-auth-azure.workspace = true litellm-auth-types.workspace = true litellm-cache.workspace = true @@ -19,6 +20,7 @@ tokio.workspace = true url.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } litellm-cache-response.workspace = true litellm-cache-testing.workspace = true rstest.workspace = true diff --git a/litellm-rust/crates/cache-azure-blob/src/cache.rs b/litellm-rust/crates/cache-azure-blob/src/cache.rs index 489b08d485e..c5c1fdd8ab9 100644 --- a/litellm-rust/crates/cache-azure-blob/src/cache.rs +++ b/litellm-rust/crates/cache-azure-blob/src/cache.rs @@ -31,7 +31,7 @@ impl AzureBlobCache { pub async fn connect( account_url: &str, container: &str, - http: reqwest::Client, + http: litellm_http::Client, codec: C, runtime: Handle, ) -> Result { diff --git a/litellm-rust/crates/cache-azure-blob/src/transport.rs b/litellm-rust/crates/cache-azure-blob/src/transport.rs index ed038b8d69d..3914b92365c 100644 --- a/litellm-rust/crates/cache-azure-blob/src/transport.rs +++ b/litellm-rust/crates/cache-azure-blob/src/transport.rs @@ -8,7 +8,7 @@ use azure_core::{ use futures_util::TryStreamExt; #[derive(Debug)] -pub struct ReqwestTransport(pub reqwest::Client); +pub struct ReqwestTransport(pub litellm_http::Client); #[async_trait::async_trait] impl HttpClient for ReqwestTransport { diff --git a/litellm-rust/crates/cache-azure-blob/tests/transport.rs b/litellm-rust/crates/cache-azure-blob/tests/transport.rs index cd1e10aa3d8..c8b14cd8543 100644 --- a/litellm-rust/crates/cache-azure-blob/tests/transport.rs +++ b/litellm-rust/crates/cache-azure-blob/tests/transport.rs @@ -31,7 +31,7 @@ async fn connect(server: &MockServer) -> AzureBlobCache> { None, ClientOptions { transport: Some(Transport::new(Arc::new(ReqwestTransport( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), )))), ..ClientOptions::default() }, diff --git a/litellm-rust/crates/cache-gcs/Cargo.toml b/litellm-rust/crates/cache-gcs/Cargo.toml index da0acf554f9..1a06683e615 100644 --- a/litellm-rust/crates/cache-gcs/Cargo.toml +++ b/litellm-rust/crates/cache-gcs/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true futures-util.workspace = true litellm-auth-gcp.workspace = true litellm-auth-types.workspace = true @@ -15,6 +16,7 @@ reqwest.workspace = true tokio.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } litellm-cache-testing.workspace = true rstest.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/cache-gcs/src/cache.rs b/litellm-rust/crates/cache-gcs/src/cache.rs index a8a7fbc9a7b..bad81573cb4 100644 --- a/litellm-rust/crates/cache-gcs/src/cache.rs +++ b/litellm-rust/crates/cache-gcs/src/cache.rs @@ -5,8 +5,8 @@ use litellm_cache::{ BaseCache, BatchCache, BatchEntry, CacheCodec, DisconnectCache, Error, ExactCacheContext, FlushCache, }; +use litellm_http::Client; use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, percent_encode}; -use reqwest::Client; use crate::{GcpTokenSource, TokenSource}; diff --git a/litellm-rust/crates/cache-gcs/tests/cache.rs b/litellm-rust/crates/cache-gcs/tests/cache.rs index cdce6a00bdd..12bb5344570 100644 --- a/litellm-rust/crates/cache-gcs/tests/cache.rs +++ b/litellm-rust/crates/cache-gcs/tests/cache.rs @@ -96,7 +96,7 @@ async fn cache_exposes_its_configuration(#[future(awt)] server: MockServer) { path_service_account: Some("/secrets/sa.json".into()), ..support::config(&server, Some("folder")) }, - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), litellm_cache::JsonCodec::::new(), ); assert_eq!(cache.bucket_name(), "bucket"); diff --git a/litellm-rust/crates/cache-gcs/tests/support/mod.rs b/litellm-rust/crates/cache-gcs/tests/support/mod.rs index 6097f0ee1bd..beb9aa39d9c 100644 --- a/litellm-rust/crates/cache-gcs/tests/support/mod.rs +++ b/litellm-rust/crates/cache-gcs/tests/support/mod.rs @@ -29,7 +29,7 @@ pub fn cache_with_token( ) -> JsonGcsCache { GcsCache::with_token_source( config(server, gcs_path), - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), JsonCodec::new(), token, ) diff --git a/litellm-rust/crates/cache-qdrant-semantic/Cargo.toml b/litellm-rust/crates/cache-qdrant-semantic/Cargo.toml index 950c2db7491..a44bef0a731 100644 --- a/litellm-rust/crates/cache-qdrant-semantic/Cargo.toml +++ b/litellm-rust/crates/cache-qdrant-semantic/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true futures-util.workspace = true litellm-cache.workspace = true qdrant-client = { workspace = true, features = ["serde"] } @@ -17,6 +18,7 @@ tokio.workspace = true uuid.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } futures-executor = "0.3" litellm-cache-testing.workspace = true rstest.workspace = true diff --git a/litellm-rust/crates/cache-qdrant-semantic/src/embedder.rs b/litellm-rust/crates/cache-qdrant-semantic/src/embedder.rs index 340393600f2..b7fbcd9b02d 100644 --- a/litellm-rust/crates/cache-qdrant-semantic/src/embedder.rs +++ b/litellm-rust/crates/cache-qdrant-semantic/src/embedder.rs @@ -1,7 +1,7 @@ use std::time::Duration; use litellm_cache::{Error, semantic::Embedder}; -use reqwest::Client; +use litellm_http::Client; use serde_json::Value; pub struct OpenAiEmbedder { diff --git a/litellm-rust/crates/cache-qdrant-semantic/tests/embedder.rs b/litellm-rust/crates/cache-qdrant-semantic/tests/embedder.rs index de0fab0a66f..24e6e5eba3e 100644 --- a/litellm-rust/crates/cache-qdrant-semantic/tests/embedder.rs +++ b/litellm-rust/crates/cache-qdrant-semantic/tests/embedder.rs @@ -5,6 +5,10 @@ use std::{ use litellm_cache::{Error, semantic::Embedder}; use litellm_cache_qdrant_semantic::{OpenAiEmbedder, OpenAiEmbedderConfig}; +use litellm_http::{ + ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, Resolution, + media::PublicDnsResolver, +}; use rstest::rstest; use serde_json::{Value, json}; use tokio::{ @@ -104,7 +108,7 @@ fn config(base: String, timeout: Option) -> OpenAiEmbedderConfig { async fn posts_embeddings_request_and_parses_vector() { let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await; let embedder = OpenAiEmbedder::new( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), config( format!("{}/", server.base_url()), Some(Duration::from_secs(1)), @@ -156,14 +160,17 @@ async fn status_timeout_and_body_errors_are_unavailable( ) { let server = TestHttpServer::response_after(status, body, Duration::from_millis(delay_ms)).await; - let embedder = OpenAiEmbedder::new(reqwest::Client::new(), config(server.base_url(), timeout)); + let embedder = OpenAiEmbedder::new( + litellm_http::Client::plain_for_test(), + config(server.base_url(), timeout), + ); assert_eq!(embedder.async_embed("hello", None).await, expected); } #[rstest] fn sync_embedding_is_unsupported() { let embedder = OpenAiEmbedder::new( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), config("http://127.0.0.1:9".to_owned(), None), ); assert_eq!( @@ -176,9 +183,12 @@ fn sync_embedding_is_unsupported() { #[tokio::test] async fn uses_the_injected_client() { let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await; - let client = reqwest::Client::builder() - .user_agent("litellm-embedder-test") - .build() + let config_with_agent = HttpClientConfig { + user_agent: Some("litellm-embedder-test".into()), + ..Resolution::from(&HttpSettings::default()).config + }; + let client = HttpClientPool::new(Arc::new(PublicDnsResolver)) + .client(&config_with_agent, ClientVariant::Provider) .unwrap(); let embedder = OpenAiEmbedder::new(client, config(server.base_url(), None)); assert_eq!( diff --git a/litellm-rust/crates/cache-s3/Cargo.toml b/litellm-rust/crates/cache-s3/Cargo.toml index c8150180e7c..680f2da8215 100644 --- a/litellm-rust/crates/cache-s3/Cargo.toml +++ b/litellm-rust/crates/cache-s3/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true litellm-cache.workspace = true litellm-auth-aws.workspace = true aws-sdk-s3 = { version = "1.146.1", default-features = false, features = ["rustls", "rt-tokio"] } @@ -19,6 +20,7 @@ reqwest.workspace = true tokio.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } litellm-cache-testing.workspace = true rstest.workspace = true wiremock = "0.6.5" diff --git a/litellm-rust/crates/cache-s3/src/auth.rs b/litellm-rust/crates/cache-s3/src/auth.rs index fdf71fc011b..f4f06e5371d 100644 --- a/litellm-rust/crates/cache-s3/src/auth.rs +++ b/litellm-rust/crates/cache-s3/src/auth.rs @@ -2,10 +2,11 @@ use aws_credential_types::{ Credentials as AwsCredentials, provider::{ProvideCredentials, error::CredentialsError, future}, }; -use litellm_auth_aws::{AwsAuthConfig, resolve_credentials}; +use litellm_auth_aws::{AwsAuthConfig, AwsAuthService}; #[derive(Clone)] pub struct S3Credentials { + auth: AwsAuthService, config: AwsAuthConfig, env: fn(&str) -> Option, } @@ -16,7 +17,11 @@ impl S3Credentials { } pub fn with_env(config: AwsAuthConfig, env: fn(&str) -> Option) -> Self { - Self { config, env } + Self { + auth: AwsAuthService::default(), + config, + env, + } } } @@ -38,7 +43,8 @@ impl ProvideCredentials for S3Credentials { "litellm-s3-cache", )); } - resolve_credentials(self.config.clone(), &self.env) + self.auth + .resolve_credentials(self.config.clone(), &self.env) .await .map_err(|_| CredentialsError::provider_error("S3 cache authentication failed")) }) diff --git a/litellm-rust/crates/cache-s3/src/cache.rs b/litellm-rust/crates/cache-s3/src/cache.rs index 91cd5e8ef54..ced948e80c6 100644 --- a/litellm-rust/crates/cache-s3/src/cache.rs +++ b/litellm-rust/crates/cache-s3/src/cache.rs @@ -42,7 +42,12 @@ pub struct S3Cache { } impl S3Cache { - pub fn new(config: S3CacheConfig, http: reqwest::Client, codec: C, runtime: Handle) -> Self { + pub fn new( + config: S3CacheConfig, + http: litellm_http::Client, + codec: C, + runtime: Handle, + ) -> Self { let endpoint_url: Option = config.endpoint.map(|endpoint| endpoint.url); let base = aws_sdk_s3::Config::builder() .behavior_version(BehaviorVersion::latest()) diff --git a/litellm-rust/crates/cache-s3/src/transport.rs b/litellm-rust/crates/cache-s3/src/transport.rs index 3e5ce578c31..eabc54ac9e2 100644 --- a/litellm-rust/crates/cache-s3/src/transport.rs +++ b/litellm-rust/crates/cache-s3/src/transport.rs @@ -9,7 +9,7 @@ use aws_smithy_runtime_api::client::{ use aws_smithy_types::body::SdkBody; #[derive(Clone, Debug)] -pub(crate) struct ReqwestHttpClient(pub(crate) reqwest::Client); +pub(crate) struct ReqwestHttpClient(pub(crate) litellm_http::Client); impl HttpClient for ReqwestHttpClient { fn http_connector( diff --git a/litellm-rust/crates/cache-s3/tests/support/mod.rs b/litellm-rust/crates/cache-s3/tests/support/mod.rs index 046b042c66a..b628c1df431 100644 --- a/litellm-rust/crates/cache-s3/tests/support/mod.rs +++ b/litellm-rust/crates/cache-s3/tests/support/mod.rs @@ -32,7 +32,12 @@ pub fn config(endpoint: &str) -> S3CacheConfig { } pub fn cache_with(config: S3CacheConfig, runtime: Handle) -> JsonS3Cache { - S3Cache::new(config, reqwest::Client::new(), JsonCodec::new(), runtime) + S3Cache::new( + config, + litellm_http::Client::plain_for_test(), + JsonCodec::new(), + runtime, + ) } pub fn cache(endpoint: &str) -> JsonS3Cache { diff --git a/litellm-rust/crates/callbacks-legacy-python/AGENTS.md b/litellm-rust/crates/callbacks-legacy-python/AGENTS.md index 8b2e1c15f6e..de6e0c1b225 100644 --- a/litellm-rust/crates/callbacks-legacy-python/AGENTS.md +++ b/litellm-rust/crates/callbacks-legacy-python/AGENTS.md @@ -1,19 +1,19 @@ - Target invariants, not completion claims -- Keep this crate the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on) plus the kwargs rewrites the wrapper makes on the way in (credential-name inheritance, the budget and retry-count limits) +- This crate is the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on) + - Smell test: if a future callback host (`callbacks-v1-python`, WASM, in-process Rust) could share a piece of this crate, it does not belong here + - SDK request policy (credential inheritance, the budget and retry-count limits) is the driver's preflight, supplied by `python-bridge`; this crate only adopts the keyword view it produces - The driver in `litellm-host-python`, the routes and core see one `PythonLifecycle`; they never learn which Python objects consume a call -- Rust drives the call; every litellm Python internal it still borrows is a variant of `LegacyPython`, grouped by subsystem (`Wrapper`, `Logging`, `DeploymentHooks`) +- Every litellm Python internal Rust still borrows is a variant of `LegacyPython`, grouped by subsystem, with its signature pinned in `python_contract.json` - The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it - Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython` - - `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the legacy path rewrites it (setup, deployment hook, prepare) and the bound request object whose attributes back keywords the caller omitted; routes hand it over through `run_legacy_call` and keep no copy -- `setup` reuses a `Logging` the caller passed as `litellm_logging_obj` (the proxy and Router are the live cases) and otherwise builds one through `function_setup`, as `@client` does - - Either way every phase calls the same `Logging` method the Python path calls; which callbacks run is `Logging`'s decision, never this crate's +- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; routes hand it over through `run_legacy_call` and keep no copy +- `setup` reuses a `Logging` passed as `litellm_logging_obj` (the proxy and Router) and otherwise builds one through `function_setup`; which callbacks run is `Logging`'s decision, never this crate's - Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation - - Retain complete boundary arguments, opaque unknown values, aliases, omitted/default distinctions and deliberate copies; preserve the established deployment-hook kwargs view - - Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object; this crate compares the two itself, and the argument is resolved by `litellm_host_python::lookup` - - Retain independently captured body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only - - A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-host`, `litellm-host-python` and the bridge; the only fact that crosses from the route is the prepared keyword view -- Success and failure handlers receive the exact selected public response or exception; logging projections, redaction and snapshots keep their own copy contracts - - Ordinary failure-handler errors cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch - - Dispatch errors never replay provider work or trigger the opposite outcome; the proxy's acceptance or rejection releases deferred success at most once + - Retain complete boundary arguments, opaque values, aliases, omitted/default distinctions and deliberate copies; preserve the deployment-hook kwargs view + - Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object, resolved through `litellm_host_python::lookup` + - Retain body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only +- Success and failure handlers receive the exact selected public response or exception + - A failure-handler error cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch + - Dispatch errors never replay provider work or trigger the opposite outcome; the proxy releases deferred success at most once - Delivery follows the registry, not the callable's type: direct, awaited, executor-submitted, logging-worker and deferred paths stay distinct - Traverse every retained Python edge; `close` is idempotent and restores the correlation context once diff --git a/litellm-rust/crates/callbacks-legacy-python/python_contract.json b/litellm-rust/crates/callbacks-legacy-python/python_contract.json index 8a7f3b98f47..9ed13ae5ed5 100644 --- a/litellm-rust/crates/callbacks-legacy-python/python_contract.json +++ b/litellm-rust/crates/callbacks-legacy-python/python_contract.json @@ -6,9 +6,6 @@ "start_time", "asynchronous" ], - "check_limits": [ - "kwargs" - ], "finalize": [ "response", "logger", @@ -76,11 +73,6 @@ ], "custom_pricing_fields": [], "is_internal_call": [], - "credential_list": [], - "warn_unknown_credential": [ - "name", - "loaded" - ], "before_deployment_call": [ "kwargs", "call_type" diff --git a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs index 75a635e9c63..fe0f6c7dd45 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs @@ -19,7 +19,7 @@ use serde_json::Value; use crate::{ DeploymentHooks, LegacyCallbacks, PublicCall, PythonLogger, deferred::{PendingLogging, PendingSuccess}, - finalize, is_internal_call, prepare, + finalize, is_internal_call, python::Streaming, setup, }; @@ -117,9 +117,13 @@ impl LegacyLogging { }) } + /// The keyword view the rest of the call reads: a copy, so the deployment hook's own + /// dict is left as the hook returned it, carrying the logger as `@client` injects it. + /// The driver's preflight rewrites this same dict before the host projects from it. fn prepare(&mut self, py: Python<'_>) -> PyResult { - let prepared = prepare(py, self.call.kwargs().bind(py), self.logger()?)?.unbind(); - self.call.set_kwargs(prepared); + let prepared = self.call.kwargs().bind(py).copy()?; + prepared.set_item("litellm_logging_obj", self.logger()?.object(py))?; + self.call.set_kwargs(prepared.unbind()); Ok(LifecycleStep::Arguments(self.call.kwargs().clone_ref(py))) } @@ -465,6 +469,7 @@ impl PythonLifecycle for LegacyLogging { error.write_unraisable(py, None); } self.body = None; + self.headers = None; self.context = None; self.stream = None; } @@ -482,7 +487,8 @@ impl PythonLifecycle for LegacyLogging { visit.call(&stream.chunks)?; visit.call(&stream.first_chunk)?; } - visit.call(&self.body) + visit.call(&self.body)?; + visit.call(&self.headers) } } @@ -580,8 +586,6 @@ assert prepared['document'] is replacement assert prepared['pages'] is replaced_kwargs['pages'] assert prepared['litellm_logging_obj'] is logger assert 'litellm_logging_obj' not in replaced_kwargs -[checked] = [value for name, value in logger.calls if name == 'check_limits'] -assert checked is prepared ", ); }); @@ -616,8 +620,6 @@ kwargs = {'logger': logger, 'vendor_extension': opaque} &locals, c" assert prepared['vendor_extension'] is opaque -[checked] = [value for name, value in logger.calls if name == 'check_limits'] -assert checked['vendor_extension'] is opaque assert hooked == ([opaque] if asynchronous else []), hooked ", ); @@ -733,45 +735,6 @@ assert all(value is failure for name, value in logger.calls if name.endswith('_h ); }); } - - #[rstest] - #[case::synchronous(false)] - #[case::asynchronous(true)] - fn a_limit_rejected_before_the_call_surfaces_as_the_callers_error(#[case] asynchronous: bool) { - Python::initialize(); - Python::attach(|py| { - let locals = namespace( - py, - c" -class BudgetExceeded(Exception): - pass - -rejection = BudgetExceeded('over budget') - -class LimitedLogger(StubLogger): - def check_limits(self, arguments): - raise rejection - -logger = LimitedLogger() -logger.hooks = {'pre': lambda kwargs: kwargs} -kwargs = {'logger': logger} -", - ); - let mut logging = legacy_call(py, &locals, asynchronous); - let kwargs = local(&locals, "kwargs") - .cast_into::() - .unwrap() - .unbind(); - let result = logging.begin(py, kwargs, 0.0).and_then(|step| match step { - LifecycleStep::Await(_) => { - logging.resume(py, Ok(local(&locals, "kwargs").unbind())) - } - step => Ok(step), - }); - let error = result.err().unwrap(); - assert!(error.value(py).is(local(&locals, "rejection"))); - }); - } } #[cfg(test)] @@ -782,6 +745,7 @@ mod payload_tests { use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest}; use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py}; use proptest::prelude::*; + use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; use rstest::rstest; use serde_json::{Map, Value, json}; @@ -871,16 +835,7 @@ check = lambda: None headers: vec![("x-route".into(), "route".into())], body, }; - let step = logging.before_send(py, Box::new(wire), &context).unwrap(); - let raw = MachineEvent::ResponseReceived { - raw: RawResponse { - body: "raw response".into(), - }, - }; - assert!(matches!( - logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(), - LifecycleStep::Done - )); + let (_, step) = send_and_receive(py, &mut logging, wire, &context); run(py, &locals, c"check()"); let LifecycleStep::Wire(wire) = step else { panic!("before_send did not hand back the wire request"); @@ -889,6 +844,134 @@ check = lambda: None }) } + /// `before_send` over `wire`, then the provider's raw response the way the driver + /// delivers it, so `pre_call` and `post_call` have both seen the retained payload. + fn send_and_receive<'a>( + py: Python<'_>, + logging: &'a mut LegacyLogging, + wire: WireRequest, + context: &RequestContext, + ) -> (&'a mut LegacyLogging, LifecycleStep) { + let step = logging.before_send(py, Box::new(wire), context).unwrap(); + let raw = MachineEvent::ResponseReceived { + raw: RawResponse { + body: "raw response".into(), + }, + }; + assert!(matches!( + logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(), + LifecycleStep::Done + )); + (logging, step) + } + + fn route_context() -> RequestContext { + RequestContext { + model: "model".into(), + custom_llm_provider: "provider".into(), + optional_params: json!({}), + secret_fields: vec![], + api_key: Some(SecretValue::new("route-key")), + } + } + + fn route_wire() -> WireRequest { + WireRequest { + url: "https://provider.invalid/ocr".into(), + headers: vec![("x-route".into(), "route".into())], + body: json!({}), + } + } + + /// A Python object owning one `LegacyLogging`, so the interpreter's collector sees the + /// edges the adapter reports and clears them the way the driver's `Execution` does. + #[pyclass(weakref)] + struct Retained { + logging: Option, + } + + #[pymethods] + impl Retained { + fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { + match &self.logging { + Some(logging) => logging.traverse(&visit), + None => Ok(()), + } + } + + fn __clear__(slf: &Bound<'_, Self>) { + drop(slf.borrow_mut().logging.take()); + } + } + + #[test] + fn a_cycle_through_the_retained_headers_is_collected() { + Python::initialize(); + Python::attach(|py| { + let locals = namespace(py, PAYLOAD_LOGGER); + let mut logging = LegacyLogging { + logger: Some(PythonLogger::new(local(&locals, "logger").unbind())), + ..legacy_call(py, &locals, false) + }; + send_and_receive(py, &mut logging, route_wire(), &route_context()); + let retained = Py::new( + py, + Retained { + logging: Some(logging), + }, + ) + .unwrap(); + locals.set_item("retained", retained).unwrap(); + run( + py, + &locals, + c" +import gc +import weakref + +logger.post[2]['headers']['owner'] = retained +logger.pre = logger.post = None +reference = weakref.ref(retained) +del retained +gc.collect() +assert reference() is None +", + ); + }); + } + + #[test] + fn close_releases_the_retained_headers() { + Python::initialize(); + Python::attach(|py| { + let locals = namespace(py, PAYLOAD_LOGGER); + let mut logging = LegacyLogging { + logger: Some(PythonLogger::new(local(&locals, "logger").unbind())), + ..legacy_call(py, &locals, false) + }; + send_and_receive(py, &mut logging, route_wire(), &route_context()); + run( + py, + &locals, + c" +import weakref + +class Sentinel: + pass + +sentinel = Sentinel() +logger.post[2]['headers']['sentinel'] = sentinel +logger.pre = logger.post = None +reference = weakref.ref(sentinel) +del sentinel +assert reference() is not None +", + ); + logging.close(py); + run(py, &locals, c"assert reference() is None"); + }); + } + #[rstest] #[case::caller_keyword(c" document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'} diff --git a/litellm-rust/crates/callbacks-legacy-python/src/call.rs b/litellm-rust/crates/callbacks-legacy-python/src/call.rs index 9b921070839..3fa638ac6d3 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/call.rs @@ -4,7 +4,7 @@ //! this crate holds them. use litellm_host::{machine::Machine, protocol::Protocol}; -use litellm_host_python::{ProtocolHost, lookup, run_call}; +use litellm_host_python::{Preflight, ProtocolHost, lookup, run_call}; use pyo3::{ gc::{PyTraverseError, PyVisit}, prelude::*, @@ -39,7 +39,8 @@ impl PublicCall { } /// The keyword view the legacy path currently reads: the caller's copy until - /// `function_setup`, then each rewrite (setup, deployment hook, prepare) in turn. + /// `function_setup`, then each rewrite (setup, deployment hook, the driver's preflight) + /// in turn. pub(crate) fn kwargs(&self) -> &Py { &self.kwargs } @@ -64,13 +65,15 @@ impl PublicCall { } /// Runs one native call under the legacy `Logging` contract: the protocol host projects from -/// the keyword view the contract prepares, and the contract observes the call. +/// the keyword view the contract prepares and `preflight` rewrites, and the contract +/// observes the call. pub fn run_legacy_call( py: Python<'_>, surface: LegacySurface, call: PublicCall, machine: M, host: H, + preflight: Preflight, asynchronous: bool, ) -> PyResult> where @@ -83,6 +86,7 @@ where machine, host, Box::new(LegacyLogging::new(py, surface, call, asynchronous)), + preflight, arguments, asynchronous, ) diff --git a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs index 69f72fbc177..8fca64d1b0e 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs @@ -1,9 +1,10 @@ //! The legacy `@client` wrapper as the native call sees it: litellm's `Logging` object, the -//! sync and async callback registries it fans out to, the deployment hooks, the deferred -//! proxy release, and the kwargs rewrites the wrapper makes on the way in (credential-name -//! inheritance, budget and retry-count limits). All of it sits behind one +//! sync and async callback registries it fans out to, the deployment hooks and the deferred +//! proxy release. All of it sits behind one //! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and -//! core never learn which Python object is on the other end. +//! core never learn which Python object is on the other end. The SDK's own request policy +//! (credential inheritance, the budget and retry limits) is the driver's preflight, not this +//! crate's. //! //! Legacy callbacks receive the caller's own objects and may mutate them. [`PublicCall`] //! is where those objects live, and [`run_legacy_call`] is how a route hands them over @@ -14,220 +15,12 @@ mod call; mod callbacks; mod deferred; mod logger; -mod preparation; mod python; pub(crate) use adapter::LegacyLogging; pub use adapter::{LegacySurface, PassThroughStream}; pub use call::{PublicCall, run_legacy_call}; pub(crate) use callbacks::{LegacyCallbacks, is_internal_call}; pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup}; -pub(crate) use preparation::prepare; #[cfg(test)] -mod test_support { - use std::ffi::CStr; - - use pyo3::prelude::*; - use pyo3::types::{PyDict, PyTuple}; - - use crate::{LegacyLogging, LegacySurface, PublicCall}; - - /// The parameters of every `callbacks_legacy_python` function, as the real module declares them. - /// `tests/unit/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python - /// signatures, and [`namespace`] binds every fake call against it. - pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json"); - - /// Stand-ins for `callbacks_legacy_python`, the only Python module the crate calls. Tests - /// share one interpreter and run concurrently, so each fake is installed idempotently and - /// forwards to the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`). - /// Every fake is bound against the contract first, so a call the real module would reject - /// fails here too. - const STUBS: &CStr = c" -import contextvars -import inspect -import json -import sys -import traceback -import types - -for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.callbacks_legacy_python'): - sys.modules.setdefault(name, types.ModuleType(name)) - -legacy = sys.modules['litellm.rust_bridge.callbacks_legacy_python'] -CONTRACT = json.loads(python_contract) - - -def contracted(name, fake): - signature = inspect.Signature( - [inspect.Parameter(parameter, inspect.Parameter.POSITIONAL_OR_KEYWORD) for parameter in CONTRACT[name]] - ) - - def checked(*args, **kwargs): - signature.bind(*args, **kwargs) - return fake(*args, **kwargs) - - return checked - - -if not hasattr(legacy, 'is_internal'): - legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False) - -FAKES = { - 'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace( - logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'], - kwargs=kwargs, - ), - 'check_limits': lambda arguments: arguments['logger'].check_limits(arguments), - 'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response), - 'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs( - kwargs=kwargs, - model=model, - optional_params=optional_params, - litellm_params=litellm_params, - custom_llm_provider=provider, - ), - 'pre_call': lambda logger, input, api_key, additional_args: logger.pre_call(input, api_key, additional_args), - 'post_call': lambda logger, original_response, api_key, additional_args: logger.post_call( - original_response, api_key, additional_args - ), - 'defers_async_logging': lambda logger: bool(getattr(logger, '_defer_async_logging', False)), - 'defer_success': lambda logger, pending: setattr(logger, '_native_pending_logging', pending), - 'sync_success_for_async_call': lambda logger, response, start, end: logger.handle_sync_success_callbacks_for_async_calls( - response, start, end - ), - 'failure_handler': lambda logger, error, start, end, asynchronous: ( - logger.async_failure_handler if asynchronous else logger.failure_handler - )(error, ''.join(traceback.format_exception(error)), start, end), - 'submit_success': lambda logger, response, start, end: logger.record('submit', (response, start, end)), - 'async_success_handler': lambda logger, response, start, end: logger.async_success_handler(response, start, end), - 'enqueue_logging': lambda coroutine: coroutine.enqueue(), - 'restore_context': lambda logger: logger.record('restore', None), - 'custom_pricing_fields': lambda: ('ocr_cost_per_page',), - 'is_internal_call': lambda: legacy.is_internal.get(), - 'credential_list': lambda: [], - 'warn_unknown_credential': lambda name, loaded: None, - 'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type), - 'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook( - 'success', response, call_type - ), - 'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type), - 'stream_opened': lambda logger: logger.record('stream_opened', None), - 'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record( - 'stream_success', list(chunks) - ), - 'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error), -} -assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys()) -for name, fake in FAKES.items(): - setattr(legacy, name, contracted(name, fake)) - - -unraisable = sys.modules.setdefault( - 'litellm_test_unraisable', types.ModuleType('litellm_test_unraisable') -) -if not hasattr(unraisable, 'events'): - unraisable.events = [] - sys.unraisablehook = lambda event: unraisable.events.append((event.object, event.exc_value)) - - -def unraisable_from(owner): - return [error for source, error in unraisable.events if source is owner] - - -class StubCoroutine: - def __init__(self, logger): - self.logger = logger - - def enqueue(self): - self.logger.record('enqueued', None) - self.logger.on_enqueue(self) - - def close(self): - self.logger.record('closed', None) - - -class StubLogger: - def __init__(self): - self.calls = [] - self.hooks = {} - self.on_enqueue = lambda coroutine: None - - def record(self, name, value): - self.calls.append((name, value)) - - def names(self): - return [name for name, _ in self.calls] - - def hook(self, phase, value, call_type): - self.record(phase + '_hook', call_type) - return self.hooks.get(phase, lambda value: 'awaitable')(value) - - def check_limits(self, arguments): - self.record('check_limits', arguments) - - def failure_handler(self, error, trace, start, end): - self.record('failure_handler', error) - - def async_failure_handler(self, error, trace, start, end): - self.record('async_failure_handler', error) - return 'awaitable' - - def success_handler(self, response, start, end): - self.record('success_handler', response) - - def async_success_handler(self, response, start, end): - self.record('async_success_handler', response) - return StubCoroutine(self) - - def handle_sync_success_callbacks_for_async_calls(self, response, start, end): - self.record('sync_success_for_async_call', response) - - -logger = StubLogger() -"; - - /// A namespace with the stubs, `StubLogger` and a fresh `logger`, after `script` ran in it. - pub(crate) fn namespace<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> { - let locals = PyDict::new(py); - locals.set_item("python_contract", PYTHON_CONTRACT).unwrap(); - py.run(STUBS, Some(&locals), Some(&locals)).unwrap(); - py.run(script, Some(&locals), Some(&locals)).unwrap(); - locals - } - - pub(crate) fn run(py: Python<'_>, locals: &Bound<'_, PyDict>, code: &CStr) { - py.run(code, Some(locals), Some(locals)).unwrap(); - } - - pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyAny> { - locals.get_item(name).unwrap().unwrap() - } - - /// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`). - pub(crate) fn legacy_call( - py: Python<'_>, - locals: &Bound<'_, PyDict>, - asynchronous: bool, - ) -> LegacyLogging { - let request = locals - .get_item("request") - .unwrap() - .unwrap_or_else(|| py.None().into_bound(py)); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .map(|kwargs| kwargs.cast_into::().unwrap()) - .unwrap_or_else(|| PyDict::new(py)); - let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap(); - LegacyLogging::new( - py, - LegacySurface { - call_type: "test", - input_description: "test input", - stream: None, - }, - call, - asynchronous, - ) - } -} +mod test_support; diff --git a/litellm-rust/crates/callbacks-legacy-python/src/python.rs b/litellm-rust/crates/callbacks-legacy-python/src/python.rs index cb609d52878..47331f369f5 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/python.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/python.rs @@ -19,18 +19,12 @@ pub(crate) enum LegacyPython { Streaming(Streaming), } -/// The `@client` wrapper around the call: `function_setup`, limits, credentials, -/// response metadata and the correlation context. +/// The `@client` wrapper around the call: `function_setup`, response metadata and the +/// correlation context. #[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)] pub(crate) enum Wrapper { #[strum(serialize = "setup")] Setup, - #[strum(serialize = "check_limits")] - CheckLimits, - #[strum(serialize = "credential_list")] - CredentialList, - #[strum(serialize = "warn_unknown_credential")] - WarnUnknownCredential, #[strum(serialize = "is_internal_call")] IsInternalCall, #[strum(serialize = "finalize")] diff --git a/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs b/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs new file mode 100644 index 00000000000..e7973a7e1a0 --- /dev/null +++ b/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs @@ -0,0 +1,199 @@ +use std::ffi::CStr; + +use pyo3::prelude::*; +use pyo3::types::{PyDict, PyTuple}; + +use crate::{LegacyLogging, LegacySurface, PublicCall}; + +/// The parameters of every `callbacks_legacy_python` function, as the real module declares them. +/// `tests/unit/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python +/// signatures, and [`namespace`] binds every fake call against it. +pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json"); + +/// Stand-ins for `callbacks_legacy_python`, the only Python module the crate calls. Tests +/// share one interpreter and run concurrently, so each fake is installed idempotently and +/// forwards to the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`). +/// Every fake is bound against the contract first, so a call the real module would reject +/// fails here too. +const STUBS: &CStr = c" +import contextvars +import inspect +import json +import sys +import traceback +import types + +for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.callbacks_legacy_python'): + sys.modules.setdefault(name, types.ModuleType(name)) + +legacy = sys.modules['litellm.rust_bridge.callbacks_legacy_python'] +CONTRACT = json.loads(python_contract) + + +def contracted(name, fake): + signature = inspect.Signature( + [inspect.Parameter(parameter, inspect.Parameter.POSITIONAL_OR_KEYWORD) for parameter in CONTRACT[name]] + ) + + def checked(*args, **kwargs): + signature.bind(*args, **kwargs) + return fake(*args, **kwargs) + + return checked + + +if not hasattr(legacy, 'is_internal'): + legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False) + +FAKES = { + 'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace( + logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'], + kwargs=kwargs, + ), + 'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response), + 'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs( + kwargs=kwargs, + model=model, + optional_params=optional_params, + litellm_params=litellm_params, + custom_llm_provider=provider, + ), + 'pre_call': lambda logger, input, api_key, additional_args: logger.pre_call(input, api_key, additional_args), + 'post_call': lambda logger, original_response, api_key, additional_args: logger.post_call( + original_response, api_key, additional_args + ), + 'defers_async_logging': lambda logger: bool(getattr(logger, '_defer_async_logging', False)), + 'defer_success': lambda logger, pending: setattr(logger, '_native_pending_logging', pending), + 'sync_success_for_async_call': lambda logger, response, start, end: logger.handle_sync_success_callbacks_for_async_calls( + response, start, end + ), + 'failure_handler': lambda logger, error, start, end, asynchronous: ( + logger.async_failure_handler if asynchronous else logger.failure_handler + )(error, ''.join(traceback.format_exception(error)), start, end), + 'submit_success': lambda logger, response, start, end: logger.record('submit', (response, start, end)), + 'async_success_handler': lambda logger, response, start, end: logger.async_success_handler(response, start, end), + 'enqueue_logging': lambda coroutine: coroutine.enqueue(), + 'restore_context': lambda logger: logger.record('restore', None), + 'custom_pricing_fields': lambda: ('ocr_cost_per_page',), + 'is_internal_call': lambda: legacy.is_internal.get(), + 'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type), + 'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook( + 'success', response, call_type + ), + 'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type), + 'stream_opened': lambda logger: logger.record('stream_opened', None), + 'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record( + 'stream_success', list(chunks) + ), + 'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error), +} +assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys()) +for name, fake in FAKES.items(): + setattr(legacy, name, contracted(name, fake)) + + +unraisable = sys.modules.setdefault( + 'litellm_test_unraisable', types.ModuleType('litellm_test_unraisable') +) +if not hasattr(unraisable, 'events'): + unraisable.events = [] + sys.unraisablehook = lambda event: unraisable.events.append((event.object, event.exc_value)) + + +def unraisable_from(owner): + return [error for source, error in unraisable.events if source is owner] + + +class StubCoroutine: + def __init__(self, logger): + self.logger = logger + + def enqueue(self): + self.logger.record('enqueued', None) + self.logger.on_enqueue(self) + + def close(self): + self.logger.record('closed', None) + + +class StubLogger: + def __init__(self): + self.calls = [] + self.hooks = {} + self.on_enqueue = lambda coroutine: None + + def record(self, name, value): + self.calls.append((name, value)) + + def names(self): + return [name for name, _ in self.calls] + + def hook(self, phase, value, call_type): + self.record(phase + '_hook', call_type) + return self.hooks.get(phase, lambda value: 'awaitable')(value) + + def failure_handler(self, error, trace, start, end): + self.record('failure_handler', error) + + def async_failure_handler(self, error, trace, start, end): + self.record('async_failure_handler', error) + return 'awaitable' + + def success_handler(self, response, start, end): + self.record('success_handler', response) + + def async_success_handler(self, response, start, end): + self.record('async_success_handler', response) + return StubCoroutine(self) + + def handle_sync_success_callbacks_for_async_calls(self, response, start, end): + self.record('sync_success_for_async_call', response) + + +logger = StubLogger() +"; + +/// A namespace with the stubs, `StubLogger` and a fresh `logger`, after `script` ran in it. +pub(crate) fn namespace<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> { + let locals = PyDict::new(py); + locals.set_item("python_contract", PYTHON_CONTRACT).unwrap(); + py.run(STUBS, Some(&locals), Some(&locals)).unwrap(); + py.run(script, Some(&locals), Some(&locals)).unwrap(); + locals +} + +pub(crate) fn run(py: Python<'_>, locals: &Bound<'_, PyDict>, code: &CStr) { + py.run(code, Some(locals), Some(locals)).unwrap(); +} + +pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyAny> { + locals.get_item(name).unwrap().unwrap() +} + +/// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`). +pub(crate) fn legacy_call( + py: Python<'_>, + locals: &Bound<'_, PyDict>, + asynchronous: bool, +) -> LegacyLogging { + let request = locals + .get_item("request") + .unwrap() + .unwrap_or_else(|| py.None().into_bound(py)); + let kwargs = locals + .get_item("kwargs") + .unwrap() + .map(|kwargs| kwargs.cast_into::().unwrap()) + .unwrap_or_else(|| PyDict::new(py)); + let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap(); + LegacyLogging::new( + py, + LegacySurface { + call_type: "test", + input_description: "test input", + stream: None, + }, + call, + asynchronous, + ) +} diff --git a/litellm-rust/crates/config/Cargo.toml b/litellm-rust/crates/config/Cargo.toml new file mode 100644 index 00000000000..36bd68fe2a0 --- /dev/null +++ b/litellm-rust/crates/config/Cargo.toml @@ -0,0 +1,16 @@ +[package] +name = "litellm-config" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-auth-types.workspace = true +serde.workspace = true +serde_yaml_ng = "0.10.0" +thiserror.workspace = true + +[dev-dependencies] +rstest.workspace = true +tempfile.workspace = true diff --git a/litellm-rust/crates/config/src/error.rs b/litellm-rust/crates/config/src/error.rs new file mode 100644 index 00000000000..61d19491abc --- /dev/null +++ b/litellm-rust/crates/config/src/error.rs @@ -0,0 +1,7 @@ +#[derive(Debug, thiserror::Error)] +pub enum Error { + #[error("could not read config")] + Read(#[from] std::io::Error), + #[error("invalid YAML config")] + Parse(#[from] serde_yaml_ng::Error), +} diff --git a/litellm-rust/crates/config/src/lib.rs b/litellm-rust/crates/config/src/lib.rs new file mode 100644 index 00000000000..8e86e345025 --- /dev/null +++ b/litellm-rust/crates/config/src/lib.rs @@ -0,0 +1,48 @@ +mod error; + +use std::path::Path; + +use litellm_auth_types::SecretValue; +use serde::Deserialize; + +pub use error::Error; + +#[derive(Clone, Debug, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct Config { + pub model_list: Box<[Model]>, + #[serde(default)] + pub general_settings: GeneralSettings, +} + +#[derive(Clone, Debug, Default, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct GeneralSettings { + pub master_key: Option, +} + +#[derive(Clone, Debug, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct Model { + pub model_name: String, + pub litellm_params: LiteLlmParams, +} + +#[derive(Clone, Debug, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct LiteLlmParams { + pub model: String, + pub api_key: Option, + pub api_base: Option, + pub custom_llm_provider: Option, +} + +impl Config { + pub fn from_yaml(yaml: &str) -> Result { + Ok(serde_yaml_ng::from_str(yaml)?) + } + + pub fn load(path: impl AsRef) -> Result { + Self::from_yaml(&std::fs::read_to_string(path)?) + } +} diff --git a/litellm-rust/crates/config/tests/config.rs b/litellm-rust/crates/config/tests/config.rs new file mode 100644 index 00000000000..ce6d684ec72 --- /dev/null +++ b/litellm-rust/crates/config/tests/config.rs @@ -0,0 +1,119 @@ +use litellm_config::{Config, Error}; +use rstest::{fixture, rstest}; +use tempfile::TempDir; + +#[fixture] +fn directory() -> TempDir { + tempfile::tempdir().unwrap() +} + +#[fixture] +fn model_list_yaml() -> &'static str { + r#" +model_list: + - model_name: assistant + litellm_params: + model: anthropic/test-model + api_key: os.environ/ANTHROPIC_API_KEY + - model_name: local + litellm_params: + model: test-model + api_base: http://localhost:8000/v1 + custom_llm_provider: openai +"# +} + +#[rstest] +fn loads_model_list_from_file(directory: TempDir, model_list_yaml: &str) { + let path = directory.path().join("config.yaml"); + std::fs::write(&path, model_list_yaml).unwrap(); + + let config = Config::load(path).unwrap(); + assert_eq!(config.model_list.len(), 2); + let anthropic = &config.model_list[0]; + assert_eq!(anthropic.model_name, "assistant"); + assert_eq!(anthropic.litellm_params.model, "anthropic/test-model"); + assert_eq!( + anthropic.litellm_params.api_key.as_ref().unwrap().expose(), + "os.environ/ANTHROPIC_API_KEY" + ); + assert!(anthropic.litellm_params.api_base.is_none()); + assert!(anthropic.litellm_params.custom_llm_provider.is_none()); + let local = &config.model_list[1]; + assert_eq!(local.model_name, "local"); + assert_eq!(local.litellm_params.model, "test-model"); + assert!(local.litellm_params.api_key.is_none()); + assert_eq!( + local.litellm_params.api_base.as_deref(), + Some("http://localhost:8000/v1") + ); + assert_eq!( + local.litellm_params.custom_llm_provider.as_deref(), + Some("openai") + ); +} + +#[rstest] +fn config_debug_redacts_api_keys() { + let config = Config::from_yaml( + "model_list: [{model_name: assistant, litellm_params: {model: anthropic/test-model, api_key: secret-value}}]", + ) + .unwrap(); + assert_eq!( + config.model_list[0] + .litellm_params + .api_key + .as_ref() + .unwrap() + .expose(), + "secret-value" + ); + assert!(!format!("{config:?}").contains("secret-value")); +} + +#[rstest] +#[case::malformed_yaml("model_list: [")] +#[case::missing_model_list("{}")] +#[case::missing_params("model_list: [{model_name: assistant}]")] +#[case::missing_model("model_list: [{model_name: assistant, litellm_params: {api_key: key}}]")] +#[case::unsupported_settings("model_list: []\ngeneral_settings: {unknown: true}")] +#[case::misspelled_param( + "model_list: [{model_name: assistant, litellm_params: {model: test, api_bsae: url}}]" +)] +fn rejects_malformed_incomplete_and_unsupported_config(#[case] yaml: &str) { + assert!(matches!(Config::from_yaml(yaml), Err(Error::Parse(_)))); +} + +#[rstest] +fn distinguishes_read_errors_from_parse_errors(directory: TempDir) { + assert!(matches!( + Config::load(directory.path().join("missing.yaml")), + Err(Error::Read(error)) if error.kind() == std::io::ErrorKind::NotFound + )); +} + +#[rstest] +#[case::literal("secret-master-key")] +#[case::reference("os.environ/LITELLM_MASTER_KEY")] +fn loads_and_redacts_the_master_key(#[case] key: &str) { + let config = Config::from_yaml(&format!( + "model_list: []\ngeneral_settings:\n master_key: {key}\n" + )) + .unwrap(); + assert_eq!( + config + .general_settings + .master_key + .as_ref() + .unwrap() + .expose(), + key + ); + assert!(!format!("{config:?}").contains(key)); +} + +#[rstest] +fn missing_general_settings_has_no_master_key() { + let config = Config::from_yaml("model_list: []").unwrap(); + assert!(config.general_settings.master_key.is_none()); +} diff --git a/litellm-rust/crates/core-utils/Cargo.toml b/litellm-rust/crates/core-utils/Cargo.toml index baf5dd16707..22196979781 100644 --- a/litellm-rust/crates/core-utils/Cargo.toml +++ b/litellm-rust/crates/core-utils/Cargo.toml @@ -13,6 +13,7 @@ serde.workspace = true serde_json.workspace = true serde_path_to_error = "0.1" serde_with.workspace = true +strum.workspace = true thiserror.workspace = true url.workspace = true diff --git a/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs b/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs index 2c4921d26be..63ef79c0fa2 100644 --- a/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs +++ b/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs @@ -11,11 +11,13 @@ //! accepts; anything richer is declined upstream by the capability gate. use litellm_types::llms::openai::{ChatMessage, ChatMessageContent}; +use strum::IntoStaticStr; pub const EMPTY_TEXT_PLACEHOLDER: &str = "[System: Empty message content sanitised to satisfy protocol]"; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq)] +#[strum(serialize_all = "snake_case")] pub enum TurnRole { User, Assistant, @@ -23,10 +25,7 @@ pub enum TurnRole { impl TurnRole { pub fn as_str(self) -> &'static str { - match self { - Self::User => "user", - Self::Assistant => "assistant", - } + self.into() } } diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 0c8a747019d..20da74789bb 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -1,4 +1,6 @@ -litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src//` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back. +litellm-core is the LiteLLM SDK in Rust. Each top-level call is a module under `src//` exposing a public entrypoint named after the route. `messages::messages()` returns `MessagesResponse::Message` for a completed response or `MessagesResponse::Stream { headers, chunks }` when the request sets `stream: true`. The chunks are Anthropic SSE bytes in a `Stream>`. Dropping the stream cancels the call. The Python bridge drives `messages::route::messages_machine()` instead, because Python has to answer the call's operations on its own thread; the gateway and the Rust SDK call the plain entrypoint + +A route module has the same five pieces, in the order Python runs them. `types.rs` holds the call, the provider request, and the response. `prepare.rs` resolves the provider and credentials and shapes the request (Python's `validate_environment`, `get_complete_url`, `transform_request`). `handler.rs` resolves auth, offers the wire request to `litellm_host::hooks::RouteHooks::before_send`, sends it, reports the raw response through `emit`, and normalizes the response or stream (`pre_call`, `post`, `post_call`, `transform_response`). `mod.rs` exposes the entrypoint that runs prepare then handler with no hooks (`()`). `route.rs`, where a host needs it, wraps the same two calls in a `CallMachine` whose `HostChannel` is the hooks, and pumps a stream through `open` and `deliver`. A handler takes `&impl RouteHooks` and never a `HostChannel` directly, so it runs without a coroutine. Keep provider transport and transformation details out of the machine driver ## Crate layering @@ -10,6 +12,16 @@ Each crate mirrors one top-level Python package, so a Rust path reads as its Pyt - `litellm-llms` mirrors `litellm/llms/`: `base_llm//transformation.rs`, `//transformation.rs`, and `base_llm/ocr/handler.rs` (the OCR request handler) - `litellm-core` mirrors the route packages (`litellm/ocr/`, `litellm/messages/`, ...): entrypoints, route request types, provider dispatch, the route machine, and hooks -A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, which each route implements over its host. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate +A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, the provider-level hooks OCR implements over its host until it folds into `litellm_host::hooks::RouteHooks`. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate + +## Error placement + +The workspace `Error definitions` rules shape each crate's error; this section decides which crate and module a failure belongs to + +A failure is declared once, by the lowest crate that raises it. Every crate above nests that error unchanged (`#[error(transparent)] Auth(#[from] litellm_auth::Error)`) or maps it once at its boundary, as `src/error.rs` does for `litellm_llms::Error`. `RouteError` collects route failures and never re-declares a variant a lower crate raises + +Scope follows the concept, not the first caller. An error type under `litellm-llms`'s `/` is private to that provider: no other provider and nothing in `base_llm` may import it. A failure two providers or two routes can hit, such as wire framing, stream event decoding, or a malformed provider response, belongs to the crate that owns the concept: `litellm-framing` for framing, `litellm_llms::Error` for the transformation layer + +`litellm_llms::Error` (`crates/llms/src/error.rs`) is the one transformation error for every provider and API. `base_llm/ocr/error.rs` is the recorded exception until OCR folds into it Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback execution of any kind. Core runs each route as a machine that yields host operations and call events; which integrations consume those events is the host's business. diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 12410c187e2..6904dcc023c 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -13,10 +13,11 @@ litellm-host.workspace = true bytes.workspace = true futures-util.workspace = true base64.workspace = true -litellm-auth.workspace = true +litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] } litellm-auth-aws.workspace = true litellm-http.workspace = true litellm-llms.workspace = true +litellm-tracing.workspace = true moka.workspace = true mime_guess = "2.0.5" rand.workspace = true @@ -36,6 +37,7 @@ url.workspace = true veil.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } litellm-auth-gcp.workspace = true litellm-llms = { workspace = true, features = ["test-support"] } rstest.workspace = true diff --git a/litellm-rust/crates/core/src/audio_transcription/client.rs b/litellm-rust/crates/core/src/audio_transcription/client.rs deleted file mode 100644 index 3cf131839b8..00000000000 --- a/litellm-rust/crates/core/src/audio_transcription/client.rs +++ /dev/null @@ -1,13 +0,0 @@ -use std::{sync::OnceLock, time::Duration}; - -use crate::constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS; - -pub(super) fn http_client() -> &'static reqwest::Client { - static CLIENT: OnceLock = OnceLock::new(); - CLIENT.get_or_init(|| { - reqwest::Client::builder() - .timeout(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)) - .build() - .unwrap_or_else(|_| reqwest::Client::new()) - }) -} diff --git a/litellm-rust/crates/core/src/audio_transcription/error.rs b/litellm-rust/crates/core/src/audio_transcription/error.rs deleted file mode 100644 index 81b57af2c6c..00000000000 --- a/litellm-rust/crates/core/src/audio_transcription/error.rs +++ /dev/null @@ -1,43 +0,0 @@ -use litellm_llms::base_llm::chat::transformation::Error as LlmError; - -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("expected {expected}, got {actual}")] - InvalidType { - expected: &'static str, - actual: &'static str, - }, - #[error("missing required field: {0}")] - MissingField(&'static str), - #[error("invalid provider: {0}")] - InvalidProvider(String), - #[error("invalid request: {0}")] - InvalidRequest(String), - #[error("invalid response: {0}")] - InvalidResponse(String), - #[error("unsupported by the rust path: {0}")] - Unsupported(&'static str), - #[error(transparent)] - Auth(#[from] litellm_auth::Error), - #[error(transparent)] - Transport(#[from] litellm_http::transport::Error), - #[error(transparent)] - Headers(#[from] litellm_http::request::HeaderError), - #[error(transparent)] - Http(#[from] litellm_http::Error), - #[error(transparent)] - Aws(#[from] litellm_auth_aws::Error), -} - -impl From for Error { - fn from(error: LlmError) -> Self { - match error { - LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual }, - LlmError::MissingField(field) => Self::MissingField(field), - LlmError::InvalidRequest(message) => Self::InvalidRequest(message), - LlmError::InvalidResponse(message) => Self::InvalidResponse(message), - LlmError::Unsupported(reason) => Self::Unsupported(reason), - LlmError::Auth(error) => Self::Auth(error), - } - } -} diff --git a/litellm-rust/crates/core/src/audio_transcription/handler.rs b/litellm-rust/crates/core/src/audio_transcription/handler.rs index a1862f341a5..866d08b22e8 100644 --- a/litellm-rust/crates/core/src/audio_transcription/handler.rs +++ b/litellm-rust/crates/core/src/audio_transcription/handler.rs @@ -1,22 +1,33 @@ -use litellm_http::request::truncate_error_body; +use std::time::Duration; + +use litellm_http::{Client, request::truncate_error_body}; +use litellm_llms::base_llm::auth::resolve_auth; use serde_json::Value; -use super::{Error, client::http_client}; -use crate::audio_transcription::types::ProviderAudioTranscriptionRequest; +use super::Error; +use crate::{ + audio_transcription::types::ProviderAudioTranscriptionRequest, + constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS, +}; pub async fn execute_audio_transcription_provider_call( + http: &Client, + auth: &litellm_auth::AuthServices, request: ProviderAudioTranscriptionRequest, ) -> Result { - let response = crate::outbound::outbound_request::( - &request.auth, + let env_lookup = |key: &str| std::env::var(key).ok(); + let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?; + let response = crate::outbound::outbound_request( + authenticated, request.url.clone(), - request.upstream_headers.clone(), &request.body, - request.timeout, - &request.optional_params, - ) - .await? - .send(http_client()) + Some( + request + .timeout + .unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)), + ), + )? + .send(http) .await .map_err(|error| { Error::Transport(litellm_http::transport::Error::Network(error.to_string())) diff --git a/litellm-rust/crates/core/src/audio_transcription/mod.rs b/litellm-rust/crates/core/src/audio_transcription/mod.rs index 801fd5e9673..3d329ebfc96 100644 --- a/litellm-rust/crates/core/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -1,16 +1,20 @@ -mod error; pub mod types; -pub use error::Error; -mod client; +pub use crate::error::RouteError as Error; mod handler; mod prepare; pub use handler::execute_audio_transcription_provider_call; +use litellm_http::{ClientVariant, HttpClientConfig}; pub use prepare::prepare_audio_transcription_provider_call; use serde_json::Value; use crate::audio_transcription::types::AudioTranscriptionRequest; -pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result { - execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?) - .await +pub async fn audio_transcription( + resources: &crate::resources::CoreResources, + config: &HttpClientConfig, + request: AudioTranscriptionRequest<'_>, +) -> Result { + let request = prepare_audio_transcription_provider_call(request)?; + let http = resources.pool.client(config, ClientVariant::Provider)?; + execute_audio_transcription_provider_call(&http, &resources.auth, request).await } diff --git a/litellm-rust/crates/core/src/audio_transcription/prepare.rs b/litellm-rust/crates/core/src/audio_transcription/prepare.rs index 807993c38b7..fa50c43d62d 100644 --- a/litellm-rust/crates/core/src/audio_transcription/prepare.rs +++ b/litellm-rust/crates/core/src/audio_transcription/prepare.rs @@ -1,7 +1,10 @@ use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}; -use litellm_http::request::{has_header, string_headers}; +use litellm_http::request::string_headers; use litellm_llms::{ - base_llm::audio_transcription::transformation::{BaseAudioTranscriptionConfig, RequestAuth}, + base_llm::{ + audio_transcription::transformation::BaseAudioTranscriptionConfig, + auth::{ValidatedEnvironment, with_default_headers}, + }, bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG, }; @@ -39,20 +42,13 @@ pub fn prepare_audio_transcription_provider_call( let config = provider_config(provider_info.custom_llm_provider) .ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?; let env_lookup = |key: &str| std::env::var(key).ok(); - let mut headers = string_headers("audio transcription", request.extra_headers)?; - let auth = config.auth_strategy(&model, &request.optional_params, &env_lookup)?; - match &auth { - RequestAuth::Bearer { token } if !has_header(&headers, "authorization") => { - headers.push(("Authorization".to_string(), format!("Bearer {token}"))); - } - RequestAuth::Header { name, value } if !has_header(&headers, name) => { - headers.push(((*name).to_string(), value.clone())); - } - RequestAuth::Bearer { .. } | RequestAuth::Header { .. } | RequestAuth::AwsSigV4 { .. } => {} - } - if !has_header(&headers, "content-type") { - headers.push(("Content-Type".to_string(), "application/json".to_string())); - } + let forwarded = string_headers("audio transcription", request.extra_headers)?; + let validated = + config.validate_environment(forwarded, &model, &request.optional_params, &env_lookup)?; + let environment = ValidatedEnvironment { + headers: with_default_headers(validated.headers, &[("Content-Type", "application/json")]), + auth: validated.auth, + }; let url = config.get_complete_url( request.api_base, &model, @@ -68,9 +64,7 @@ pub fn prepare_audio_transcription_provider_call( config, url, body: transformed.body, - upstream_headers: headers, - auth, - optional_params: request.optional_params, + environment, timeout: request.timeout, }) } diff --git a/litellm-rust/crates/core/src/audio_transcription/types.rs b/litellm-rust/crates/core/src/audio_transcription/types.rs index 0d87483c9bf..eff30c1e19a 100644 --- a/litellm-rust/crates/core/src/audio_transcription/types.rs +++ b/litellm-rust/crates/core/src/audio_transcription/types.rs @@ -1,7 +1,7 @@ use std::time::Duration; -use litellm_llms::base_llm::audio_transcription::transformation::{ - BaseAudioTranscriptionConfig, RequestAuth, +use litellm_llms::base_llm::{ + audio_transcription::transformation::BaseAudioTranscriptionConfig, auth::ValidatedEnvironment, }; use serde_json::{Map, Value}; @@ -23,9 +23,7 @@ pub struct ProviderAudioTranscriptionRequest { pub config: &'static dyn BaseAudioTranscriptionConfig, pub url: String, pub body: Value, - pub upstream_headers: Vec<(String, String)>, - pub auth: RequestAuth, - pub optional_params: Map, + pub environment: ValidatedEnvironment, pub timeout: Option, } diff --git a/litellm-rust/crates/core/src/chat_completions/client.rs b/litellm-rust/crates/core/src/chat_completions/client.rs deleted file mode 100644 index d8ad6c49b7b..00000000000 --- a/litellm-rust/crates/core/src/chat_completions/client.rs +++ /dev/null @@ -1,14 +0,0 @@ -use std::{sync::OnceLock, time::Duration}; - -use crate::constants::{CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS, CHAT_COMPLETIONS_TIMEOUT_SECS}; - -pub(super) fn http_client() -> &'static reqwest::Client { - static CLIENT: OnceLock = OnceLock::new(); - CLIENT.get_or_init(|| { - reqwest::Client::builder() - .timeout(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS)) - .connect_timeout(Duration::from_secs(CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS)) - .build() - .unwrap_or_else(|_| reqwest::Client::new()) - }) -} diff --git a/litellm-rust/crates/core/src/chat_completions/error.rs b/litellm-rust/crates/core/src/chat_completions/error.rs deleted file mode 100644 index 81b57af2c6c..00000000000 --- a/litellm-rust/crates/core/src/chat_completions/error.rs +++ /dev/null @@ -1,43 +0,0 @@ -use litellm_llms::base_llm::chat::transformation::Error as LlmError; - -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("expected {expected}, got {actual}")] - InvalidType { - expected: &'static str, - actual: &'static str, - }, - #[error("missing required field: {0}")] - MissingField(&'static str), - #[error("invalid provider: {0}")] - InvalidProvider(String), - #[error("invalid request: {0}")] - InvalidRequest(String), - #[error("invalid response: {0}")] - InvalidResponse(String), - #[error("unsupported by the rust path: {0}")] - Unsupported(&'static str), - #[error(transparent)] - Auth(#[from] litellm_auth::Error), - #[error(transparent)] - Transport(#[from] litellm_http::transport::Error), - #[error(transparent)] - Headers(#[from] litellm_http::request::HeaderError), - #[error(transparent)] - Http(#[from] litellm_http::Error), - #[error(transparent)] - Aws(#[from] litellm_auth_aws::Error), -} - -impl From for Error { - fn from(error: LlmError) -> Self { - match error { - LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual }, - LlmError::MissingField(field) => Self::MissingField(field), - LlmError::InvalidRequest(message) => Self::InvalidRequest(message), - LlmError::InvalidResponse(message) => Self::InvalidResponse(message), - LlmError::Unsupported(reason) => Self::Unsupported(reason), - LlmError::Auth(error) => Self::Auth(error), - } - } -} diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index 2391ab83a60..740db2edefb 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -1,20 +1,70 @@ -use litellm_http::{outbound::OutboundRequest, request::truncate_error_body}; -use litellm_llms::base_llm::chat::transformation::ProviderChatResponseData; +use std::time::Duration; + +use litellm_auth::AuthServices; +use litellm_host::{ + event::{MachineEvent, RawResponse, RequestContext, WireRequest}, + hooks::RouteHooks, +}; +use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body}; +use litellm_llms::base_llm::{ + auth::{Authenticated, resolve_auth}, + chat::transformation::ProviderChatResponseData, +}; use litellm_types::utils::ChatCompletionsResponse; use serde_json::Value; -use super::{Error, client::http_client, prepare::prepare_provider_request}; -use crate::chat_completions::types::{ - ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest, +use super::Error; +use crate::{ + chat_completions::types::ProviderChatCompletionsRequest, + constants::CHAT_COMPLETIONS_TIMEOUT_SECS, }; -pub(super) async fn execute_chat_completions_provider_call( - request: ResolvedChatCompletionsRequest<'_>, +pub(super) async fn execute( + http: &Client, + auth: &AuthServices, + request: ProviderChatCompletionsRequest, + hooks: &impl RouteHooks, ) -> Result { - let request = prepare_provider_request(request)?; - let outbound = outbound_request(&request).await?; + let ProviderChatCompletionsRequest { + model, + custom_llm_provider, + config, + url, + body, + optional_params, + environment, + timeout, + api_key, + } = request; + let context = RequestContext { + model: model.clone(), + custom_llm_provider, + optional_params: Value::Object(optional_params), + secret_fields: Vec::new(), + api_key, + }; + let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?; + let wire = hooks + .before_send( + WireRequest { + url, + headers: authenticated.headers, + body, + }, + context, + ) + .await?; + let outbound = outbound_request( + Authenticated { + headers: wire.headers, + signer: authenticated.signer, + }, + wire.url, + &wire.body, + timeout, + )?; - let response = outbound.send(http_client()).await.map_err(|err| { + let response = outbound.send(http).await.map_err(|err| { // Failing to establish the connection means the request never went out, // so the host can still serve it. Everything else here, a timeout // above all, may have reached the provider and been answered. @@ -36,13 +86,17 @@ pub(super) async fn execute_chat_completions_provider_call( body: truncate_error_body(&text), })); } + hooks + .emit(MachineEvent::ResponseReceived { + raw: RawResponse { body: text.clone() }, + }) + .await?; let body: Value = serde_json::from_str(&text).map_err(|err| { Error::InvalidResponse(format!("invalid chat completions response JSON: {err}")) })?; - request - .config - .transform_response(&request.model, ProviderChatResponseData { body }) + config + .transform_response(&model, ProviderChatResponseData { body }) .map_err(Error::from) .map_err(as_response_error) } @@ -64,31 +118,157 @@ pub(super) fn as_response_error(err: Error) -> Error { } } -pub(super) async fn outbound_request( - request: &ProviderChatCompletionsRequest, +pub(super) fn outbound_request( + authenticated: Authenticated, + url: String, + body: &Value, + timeout: Option, ) -> Result { crate::outbound::outbound_request( - &request.auth, - request.url.clone(), - request.upstream_headers.clone(), - &request.body, - request.timeout, - &request.optional_params, + authenticated, + url, + body, + Some(timeout.unwrap_or(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS))), ) - .await .map_err(|error| match error { // Python drops the caller's copy and prefers a forwarded Authorization // over the signature, so leave the request to it. - Error::Http(litellm_http::Error::ComputedHeader(_)) => { + litellm_http::Error::ComputedHeader(_) => { Error::Unsupported("request forwards a header AWS SigV4 computes") } - other => other, + other => Error::Http(other), }) } #[cfg(test)] mod tests { - use super::{Error, as_response_error}; + use std::sync::Mutex; + + use rstest::rstest; + use serde_json::json; + use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any}; + + use super::*; + use crate::chat_completions::{ + prepare::{prepare_provider_request, resolve_request}, + types::ChatCompletionsRequest, + }; + + const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#; + + /// Rewrites the outgoing request and records what the call reports back. + #[derive(Default)] + struct RecordingHooks { + contexts: Mutex>, + raw: Mutex>, + } + + impl RouteHooks for RecordingHooks { + async fn before_send( + &self, + wire: WireRequest, + context: RequestContext, + ) -> Result { + self.contexts.lock().unwrap().push(context); + let mut body = wire.body; + body["system"] = json!("added by the host"); + Ok(WireRequest { + headers: wire + .headers + .into_iter() + .chain([("x-host".to_string(), "seen".to_string())]) + .collect(), + body, + ..wire + }) + } + + async fn emit(&self, event: MachineEvent) -> Result<(), Error> { + let MachineEvent::ResponseReceived { raw } = event; + self.raw.lock().unwrap().push(raw.body); + Ok(()) + } + } + + fn prepared(api_base: &str) -> ProviderChatCompletionsRequest { + prepare_provider_request( + resolve_request(ChatCompletionsRequest { + model: "anthropic/claude-sonnet-4-5", + messages: json!([{"role": "user", "content": "hi"}]), + optional_params: json!({"max_tokens": 16}).as_object().unwrap().clone(), + api_key: Some("sk-test"), + api_base: Some(api_base), + custom_llm_provider: None, + extra_headers: None, + timeout: None, + }) + .unwrap(), + ) + .unwrap() + } + + #[rstest] + #[tokio::test] + async fn the_hooks_rewrite_the_wire_request_and_see_the_raw_response() { + let upstream = MockServer::start().await; + Mock::given(any()) + .respond_with( + ResponseTemplate::new(200).set_body_raw(ANTHROPIC_MESSAGE, "application/json"), + ) + .mount(&upstream) + .await; + let hooks = RecordingHooks::default(); + + execute( + &Client::plain_for_test(), + &AuthServices::default(), + prepared(&upstream.uri()), + &hooks, + ) + .await + .expect("chat completions call succeeds"); + + let [request] = <[Request; 1]>::try_from(upstream.received_requests().await.unwrap()) + .unwrap_or_else(|requests| panic!("one request, saw {}", requests.len())); + let sent: Value = serde_json::from_slice(&request.body).unwrap(); + assert_eq!(sent["system"], "added by the host"); + assert_eq!(request.headers["x-host"], "seen"); + assert_eq!(request.headers["x-api-key"], "sk-test"); + let [context] = <[RequestContext; 1]>::try_from(hooks.contexts.into_inner().unwrap()) + .unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len())); + assert_eq!( + (context.model.as_str(), context.custom_llm_provider.as_str()), + ("claude-sonnet-4-5", "anthropic") + ); + assert_eq!(context.optional_params, json!({"max_tokens": 16})); + assert_eq!(hooks.raw.into_inner().unwrap(), [ANTHROPIC_MESSAGE]); + } + + #[rstest] + #[tokio::test] + async fn an_upstream_failure_is_not_reported_as_a_received_response() { + let upstream = MockServer::start().await; + Mock::given(any()) + .respond_with(ResponseTemplate::new(500).set_body_string("boom")) + .mount(&upstream) + .await; + let hooks = RecordingHooks::default(); + + let error = execute( + &Client::plain_for_test(), + &AuthServices::default(), + prepared(&upstream.uri()), + &hooks, + ) + .await + .expect_err("the upstream failure fails the call"); + + assert!(matches!( + error, + Error::Transport(litellm_http::transport::Error::Http { status: 500, .. }) + )); + assert!(hooks.raw.into_inner().unwrap().is_empty()); + } #[test] fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() { diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 224c9d8cfed..dc4e80b816a 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -6,24 +6,26 @@ //! credentials, and it resolves the provider, translates the conversation, //! calls the provider, and returns a typed OpenAI-shaped response. -mod error; pub mod types; -pub use error::Error; -mod client; +pub use crate::error::RouteError as Error; mod common_utils; pub(crate) mod handler; mod prepare; -use handler::execute_chat_completions_provider_call; +use litellm_http::{ClientVariant, HttpClientConfig}; use litellm_types::utils::ChatCompletionsResponse; -use prepare::{parse_messages, resolve_provider_config, resolve_request}; +use prepare::{parse_messages, prepare_provider_request, resolve_provider_config, resolve_request}; use serde_json::{Map, Value}; use crate::chat_completions::types::ChatCompletionsRequest; pub async fn chat_completions( + resources: &crate::resources::CoreResources, + config: &HttpClientConfig, request: ChatCompletionsRequest<'_>, ) -> Result { - execute_chat_completions_provider_call(resolve_request(request)?).await + let http = resources.pool.client(config, ClientVariant::Provider)?; + let request = prepare_provider_request(resolve_request(request)?)?; + handler::execute(&http, &resources.auth, request, &()).await } /// Whether the core would accept this request, without resolving credentials or @@ -38,9 +40,10 @@ pub fn chat_completions_decline_reason( messages: Value, optional_params: &Map, ) -> Option<&'static str> { - let Ok((_, config)) = resolve_provider_config(model, custom_llm_provider) else { + let Ok(resolved) = resolve_provider_config(model, custom_llm_provider) else { return Some("provider is not on the rust chat completions path"); }; + let config = resolved.config; let Ok(messages) = parse_messages(messages) else { return Some("unreadable message list"); }; diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index b6425773964..6ead713eec1 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -1,6 +1,9 @@ +use litellm_auth::SecretValue; use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}; -use litellm_http::request::has_header; -use litellm_llms::base_llm::chat::transformation::{BaseConfig, RequestAuth}; +use litellm_llms::base_llm::{ + auth::{ValidatedEnvironment, with_default_headers}, + chat::transformation::BaseConfig, +}; use litellm_types::llms::openai::ChatMessage; use serde_json::Value; @@ -12,10 +15,16 @@ use crate::chat_completions::types::{ ChatCompletionsRequest, ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest, }; +pub(super) struct ResolvedProvider { + pub(super) model: String, + pub(super) custom_llm_provider: String, + pub(super) config: &'static dyn BaseConfig, +} + pub(super) fn resolve_provider_config<'a>( model: &'a str, custom_llm_provider: Option<&'a str>, -) -> Result<(String, &'static dyn BaseConfig), Error> { +) -> Result { let provider_info = get_custom_llm_provider(model, custom_llm_provider) .or_else(|| { custom_llm_provider.map(|provider| CustomLlmProvider { @@ -30,7 +39,11 @@ pub(super) fn resolve_provider_config<'a>( })?; let config = chat_completions_provider_config(provider_info.custom_llm_provider) .ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?; - Ok((provider_info.model.to_string(), config)) + Ok(ResolvedProvider { + model: provider_info.model.to_string(), + custom_llm_provider: provider_info.custom_llm_provider.to_string(), + config, + }) } pub(super) fn parse_messages(messages: Value) -> Result, Error> { @@ -41,7 +54,11 @@ pub(super) fn parse_messages(messages: Value) -> Result, Error> pub(super) fn resolve_request( request: ChatCompletionsRequest<'_>, ) -> Result, Error> { - let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)?; + let ResolvedProvider { + model, + custom_llm_provider, + config, + } = resolve_provider_config(request.model, request.custom_llm_provider)?; let messages = parse_messages(request.messages)?; if messages.is_empty() { return Err(Error::InvalidRequest( @@ -53,6 +70,7 @@ pub(super) fn resolve_request( } Ok(ResolvedChatCompletionsRequest { model, + custom_llm_provider, config, messages, optional_params: request.optional_params, @@ -67,59 +85,26 @@ fn validate_environment( request: &ResolvedChatCompletionsRequest<'_>, model: &str, config: &dyn BaseConfig, -) -> Result<(Vec<(String, String)>, RequestAuth), Error> { +) -> Result { let env_lookup = |key: &str| std::env::var(key).ok(); - let mut headers = string_headers(request.extra_headers.clone())?; - let auth = config.auth( + let forwarded = string_headers(request.extra_headers.clone())?; + let validated = config.validate_environment( + forwarded, request.api_key, model, &request.optional_params, &env_lookup, )?; - match &auth { - RequestAuth::Header { name, value } => { - // The deployment's credential replaces whatever the caller forwarded - // under the same name, mirroring Python's - // `{**headers, **anthropic_headers}`: letting a request header win - // would let its sender choose the principal the call bills to. - // - // The exception is a scheme the provider hands off to entirely, such - // as an Anthropic OAuth bearer, where Python drops `x-api-key` - // instead of resolving one. Re-adding it there would put the - // credential into a header the host removed on purpose. - if !config.defers_to_forwarded_auth(&headers) { - headers.retain(|(header, _)| !header.eq_ignore_ascii_case(name)); - headers.push(((*name).to_string(), value.clone())); - } - } - RequestAuth::Bearer { token } => { - // Bedrock's `get_request_headers` assigns `headers["Authorization"]` - // unconditionally once a bearer token resolves, so the deployment's - // identity outranks whatever the caller forwarded. Keeping the - // caller's would bill and authorize the call as a different - // principal than the same deployment uses on Python. - // - // The `Header` arm below keeps the opposite precedence on purpose: - // Anthropic's transform honours a forwarded OAuth bearer. - headers.retain(|(name, _)| !name.eq_ignore_ascii_case("authorization")); - headers.push(("authorization".to_string(), format!("Bearer {token}"))); - } - // SigV4 signs the serialized body, so the handler adds its headers. - RequestAuth::AwsSigV4 { .. } => {} - } - - for (name, value) in config.default_headers() { - if !has_header(&headers, name) { - headers.push(((*name).to_string(), (*value).to_string())); - } - } - Ok((headers, auth)) + Ok(ValidatedEnvironment { + headers: with_default_headers(validated.headers, config.default_headers()), + auth: validated.auth, + }) } pub(super) fn prepare_provider_request( request: ResolvedChatCompletionsRequest<'_>, ) -> Result { - let (headers, auth) = validate_environment(&request, &request.model, request.config)?; + let environment = validate_environment(&request, &request.model, request.config)?; let model = request.model; let config = request.config; let env_lookup = |key: &str| std::env::var(key).ok(); @@ -134,19 +119,21 @@ pub(super) fn prepare_provider_request( Ok(ProviderChatCompletionsRequest { model, + custom_llm_provider: request.custom_llm_provider, config, url, body: transformed.body, - upstream_headers: headers, - auth, optional_params: request.optional_params, + environment, timeout: request.timeout, + api_key: request.api_key.map(|key| SecretValue::new(key.to_string())), }) } #[cfg(test)] mod tests { - use litellm_llms::base_llm::chat::transformation::RequestAuth; + use litellm_auth::CredentialPlacement; + use litellm_llms::base_llm::auth::{AuthScheme, resolve_auth}; use serde_json::{Map, Value, json}; use super::{prepare_provider_request, resolve_request}; @@ -161,6 +148,20 @@ mod tests { prepare_provider_request(resolve_request(request)?) } + /// The headers as they go on the wire, credential applied. + fn wire_headers(prepared: &ProviderChatCompletionsRequest) -> Vec<(String, String)> { + tokio::runtime::Builder::new_current_thread() + .build() + .unwrap() + .block_on(resolve_auth( + &litellm_auth::AuthServices::default(), + prepared.environment.clone(), + &|_| None, + )) + .unwrap() + .headers + } + fn request<'a>( model: &'a str, provider: Option<&'a str>, @@ -227,19 +228,16 @@ mod tests { )) .expect("prepares"); assert!( - prepared - .upstream_headers - .contains(&("x-api-key".to_string(), "sk-test".to_string())) + wire_headers(&prepared).contains(&("x-api-key".to_string(), "sk-test".to_string())) ); assert!( - prepared - .upstream_headers + wire_headers(&prepared) .contains(&("anthropic-version".to_string(), "2023-06-01".to_string())) ); assert!(matches!( - prepared.auth, - RequestAuth::Header { - name: "x-api-key", + prepared.environment.auth, + AuthScheme::Credential { + placement: CredentialPlacement::Header("x-api-key"), .. } )); @@ -261,12 +259,12 @@ mod tests { json!("sk-caller"), )])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let keys: Vec<_> = prepared - .upstream_headers + let headers = wire_headers(&prepared); + let keys: Vec<_> = headers .iter() .filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key")) .collect(); - assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers); + assert_eq!(keys.len(), 1, "got {:?}", headers); assert_eq!(keys[0].1, "sk-test"); } @@ -290,16 +288,14 @@ mod tests { ])); let prepared = prepare_chat_completions_call(call).expect("prepares"); assert!( - !prepared - .upstream_headers + !wire_headers(&prepared) .iter() .any(|(name, value)| name.eq_ignore_ascii_case("x-api-key") && value == "sk-test"), "the resolved key must not be applied over an OAuth bearer, got {:?}", - prepared.upstream_headers + wire_headers(&prepared) ); assert!( - prepared - .upstream_headers + wire_headers(&prepared) .iter() .any(|(name, value)| name.eq_ignore_ascii_case("authorization") && value == "Bearer sk-ant-oat01-token") @@ -322,21 +318,20 @@ mod tests { ("X-Api-Key".to_string(), json!("sk-caller")), ])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let keys: Vec<_> = prepared - .upstream_headers + let headers = wire_headers(&prepared); + let keys: Vec<_> = headers .iter() .filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key")) .collect(); - assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers); + assert_eq!(keys.len(), 1, "got {:?}", headers); assert_eq!(keys[0].1, "sk-test"); assert!( - prepared - .upstream_headers + wire_headers(&prepared) .iter() .any(|(name, value)| name.eq_ignore_ascii_case("authorization") && value == "Bearer unrelated"), "the unrelated authorization must survive, got {:?}", - prepared.upstream_headers + wire_headers(&prepared) ); } @@ -435,18 +430,16 @@ mod tests { prepared.url, "https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/converse" ); - assert_eq!( - prepared.auth, - RequestAuth::AwsSigV4 { - region: "us-east-1".to_string(), - service: "bedrock", - } - ); + assert!(matches!( + &prepared.environment.auth, + AuthScheme::AwsSigV4 { region, service: "bedrock", .. } if region == "us-east-1" + )); // SigV4 signs the serialized body, so prepare must not have added an - // Authorization header; the handler does it. + // Authorization header; the signer does it over the bytes sent. assert!( !prepared - .upstream_headers + .environment + .headers .iter() .any(|(name, _)| name.eq_ignore_ascii_case("authorization")) ); @@ -475,9 +468,20 @@ mod tests { json!("abc-123"), )])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let signed = crate::chat_completions::handler::outbound_request(&prepared) - .await - .expect("signs"); + let authenticated = resolve_auth( + &litellm_auth::AuthServices::default(), + prepared.environment, + &|_| None, + ) + .await + .expect("resolves"); + let signed = crate::chat_completions::handler::outbound_request( + authenticated, + prepared.url, + &prepared.body, + prepared.timeout, + ) + .expect("signs"); let authorization = signed .header("authorization") @@ -525,9 +529,20 @@ mod tests { call.api_key = None; call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let error = crate::chat_completions::handler::outbound_request(&prepared) - .await - .expect_err("{forwarded} should decline instead of being signed"); + let authenticated = resolve_auth( + &litellm_auth::AuthServices::default(), + prepared.environment, + &|_| None, + ) + .await + .expect("resolves"); + let error = crate::chat_completions::handler::outbound_request( + authenticated, + prepared.url, + &prepared.body, + prepared.timeout, + ) + .expect_err("{forwarded} should decline instead of being signed"); assert!( matches!(error, Error::Unsupported(_)), "{forwarded} declined as {error:?}, which the host would not fall back on" @@ -552,8 +567,8 @@ mod tests { json!("Bearer caller-supplied"), )])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let authorizations: Vec<_> = prepared - .upstream_headers + let headers = wire_headers(&prepared); + let authorizations: Vec<_> = headers .iter() .filter(|(name, _)| name.eq_ignore_ascii_case("authorization")) .map(|(_, value)| value.as_str()) @@ -585,16 +600,15 @@ mod tests { json!("Bearer sk-ant-oat01-forwarded"), )])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let keys: Vec<_> = prepared - .upstream_headers + let headers = wire_headers(&prepared); + let keys: Vec<_> = headers .iter() .filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key")) .map(|(_, value)| value.as_str()) .collect(); - assert!(keys.is_empty(), "got {:?}", prepared.upstream_headers); + assert!(keys.is_empty(), "got {:?}", headers); assert!( - prepared - .upstream_headers + wire_headers(&prepared) .iter() .any(|(name, value)| name.eq_ignore_ascii_case("authorization") && value == "Bearer sk-ant-oat01-forwarded") @@ -613,15 +627,13 @@ mod tests { json!({"maxTokens": 16}), )) .expect("prepares"); - assert_eq!( - prepared.auth, - RequestAuth::Bearer { - token: "sk-test".to_string() - } - ); + assert!(matches!( + &prepared.environment.auth, + AuthScheme::Credential { placement: CredentialPlacement::Bearer, secret } + if secret.expose() == "sk-test" + )); assert!( - prepared - .upstream_headers + wire_headers(&prepared) .iter() .any(|(name, value)| name.eq_ignore_ascii_case("authorization") && value == "Bearer sk-test"), diff --git a/litellm-rust/crates/core/src/chat_completions/types.rs b/litellm-rust/crates/core/src/chat_completions/types.rs index 3b74cf5dace..8c969dee730 100644 --- a/litellm-rust/crates/core/src/chat_completions/types.rs +++ b/litellm-rust/crates/core/src/chat_completions/types.rs @@ -1,6 +1,7 @@ use std::time::Duration; -use litellm_llms::base_llm::chat::transformation::{BaseConfig, RequestAuth}; +use litellm_auth::SecretValue; +use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig}; use litellm_types::llms::openai::ChatMessage; use serde_json::{Map, Value}; @@ -23,6 +24,7 @@ pub struct ChatCompletionsRequest<'a> { pub struct ResolvedChatCompletionsRequest<'a> { pub model: String, + pub custom_llm_provider: String, pub config: &'static dyn BaseConfig, pub messages: Vec, pub optional_params: Map, @@ -34,11 +36,16 @@ pub struct ResolvedChatCompletionsRequest<'a> { pub struct ProviderChatCompletionsRequest { pub model: String, + pub custom_llm_provider: String, pub config: &'static dyn BaseConfig, pub url: String, pub body: Value, - pub upstream_headers: Vec<(String, String)>, - pub auth: RequestAuth, + /// The route's parameters before the provider transformation, reported to the host + /// beside the wire request. pub optional_params: Map, + /// The forwarded and default headers plus how the call authenticates; the credential + /// itself is applied when the request is sent. + pub environment: ValidatedEnvironment, pub timeout: Option, + pub api_key: Option, } diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index 3d740e39677..c14b54679ff 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -5,20 +5,10 @@ pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com"; /// timeout from the caller still overrides this on the request builder. pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600; -/// Connect timeout for Anthropic Messages provider calls, in seconds. -pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10; - -/// Provider name used for Anthropic Messages when a deployment's provider model -/// does not carry an explicit provider prefix. -pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; - /// Full-request timeout ceiling for chat completions provider calls, in /// seconds. Mirrors the Python chat completions default. pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600; -/// Connect timeout for chat completions provider calls, in seconds. -pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10; - pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600; /// `object` field every non-streaming chat completion response carries. diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index eb4cd2367ec..0d3de6e57c1 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -1,15 +1,166 @@ -use litellm_llms::base_llm::ocr::error::Error as OcrError; +//! One error for every route in this crate. OCR still carries its own, richer enum. +//! +//! A variant is declared by the layer that produces it and nested here as is: +//! credentials by `litellm_auth` (AWS folds into it at that crate's boundary), the wire by +//! `litellm_http`, secrets by `litellm_secrets`. The transformation layer's [`LlmError`] +//! maps onto the same-named variants once, here, so no route re-declares them. -#[derive(Debug, thiserror::Error)] -pub enum Error { +use std::sync::Arc; + +use litellm_http::transport::Error as TransportError; +use litellm_llms::Error as LlmError; + +#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] +pub enum RouteError { + #[error("expected {expected}, got {actual}")] + InvalidType { + expected: &'static str, + actual: &'static str, + }, + #[error("missing required field: {0}")] + MissingField(&'static str), + #[error("invalid provider: {0}")] + InvalidProvider(String), + #[error("invalid request: {0}")] + InvalidRequest(String), + #[error("invalid response: {0}")] + InvalidResponse(String), + #[error("unsupported by the rust path: {0}")] + Unsupported(&'static str), #[error(transparent)] - Ocr(#[from] OcrError), + Auth(#[from] litellm_auth::Error), #[error(transparent)] - Messages(#[from] crate::messages::Error), + Transport(#[from] TransportError), #[error(transparent)] - ChatCompletions(#[from] crate::chat_completions::Error), + Headers(#[from] litellm_http::request::HeaderError), #[error(transparent)] - AudioTranscription(#[from] crate::audio_transcription::Error), + Http(#[from] litellm_http::Error), #[error(transparent)] - Responses(#[from] crate::responses::Error), + Secret(#[from] SecretError), +} + +/// Whether the provider had already been called when the route failed. Before the send, a +/// host may retry on another path; after it, the provider has done the work and billed for it. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Phase { + BeforeSend, + AfterSend, +} + +impl RouteError { + pub fn phase(&self) -> Phase { + match self { + Self::InvalidResponse(_) + | Self::Transport(TransportError::Http { .. } | TransportError::Network(_)) => { + Phase::AfterSend + } + Self::Transport(TransportError::Connect(_)) + | Self::InvalidType { .. } + | Self::MissingField(_) + | Self::InvalidProvider(_) + | Self::InvalidRequest(_) + | Self::Unsupported(_) + | Self::Auth(_) + | Self::Headers(_) + | Self::Http(_) + | Self::Secret(_) => Phase::BeforeSend, + } + } + + /// The caller's request is what is wrong, as opposed to the environment, the wire, or + /// the provider's answer. + pub fn is_request(&self) -> bool { + match self { + Self::InvalidType { .. } + | Self::MissingField(_) + | Self::InvalidProvider(_) + | Self::InvalidRequest(_) + | Self::Unsupported(_) + | Self::Headers(_) => true, + Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }), + Self::InvalidResponse(_) | Self::Transport(_) | Self::Http(_) | Self::Secret(_) => { + false + } + } + } +} + +impl From for RouteError { + fn from(error: LlmError) -> Self { + match error { + LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual }, + LlmError::MissingField(field) => Self::MissingField(field), + LlmError::InvalidRequest(message) => Self::InvalidRequest(message), + LlmError::InvalidResponse(message) => Self::InvalidResponse(message), + LlmError::Unsupported(reason) => Self::Unsupported(reason), + LlmError::Auth(error) => Self::Auth(error), + } + } +} + +#[derive(Clone, Debug, thiserror::Error)] +#[error(transparent)] +pub struct SecretError(Arc); + +impl SecretError { + pub fn source_error(&self) -> &litellm_secrets::Error { + &self.0 + } +} + +impl From for RouteError { + fn from(error: litellm_secrets::Error) -> Self { + Self::Secret(SecretError(Arc::new(error))) + } +} + +impl PartialEq for SecretError { + fn eq(&self, other: &Self) -> bool { + Arc::ptr_eq(&self.0, &other.0) + } +} + +impl Eq for SecretError {} + +#[cfg(test)] +mod tests { + use super::{Phase, RouteError}; + use litellm_http::transport::Error as TransportError; + + #[test] + fn only_a_provider_answer_or_a_lost_connection_counts_as_after_send() { + let after = [ + RouteError::InvalidResponse("bad json".into()), + RouteError::Transport(TransportError::Http { + status: 500, + body: "boom".into(), + }), + RouteError::Transport(TransportError::Network("reset".into())), + ]; + for error in after { + assert_eq!(error.phase(), Phase::AfterSend, "{error:?}"); + } + let before = [ + RouteError::Transport(TransportError::Connect("refused".into())), + RouteError::Unsupported("streaming"), + RouteError::Auth(litellm_auth::Error::InvalidHeader), + ]; + for error in before { + assert_eq!(error.phase(), Phase::BeforeSend, "{error:?}"); + } + } + + #[test] + fn a_missing_api_key_is_the_environment_not_the_request() { + assert!( + !RouteError::Auth(litellm_auth::Error::MissingApiKey { + provider: "Anthropic", + environment_variable: "ANTHROPIC_API_KEY", + }) + .is_request() + ); + assert!(RouteError::Auth(litellm_auth::Error::InvalidHeader).is_request()); + assert!(RouteError::InvalidRequest("top_k".into()).is_request()); + assert!(!RouteError::InvalidResponse("bad json".into()).is_request()); + } } diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index afe5ea595aa..d373262ae7d 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -5,6 +5,7 @@ pub mod error; pub mod messages; pub mod ocr; mod outbound; +pub mod resources; pub mod responses; -pub use error::Error; +pub use error::{Phase, RouteError}; diff --git a/litellm-rust/crates/core/src/messages/client.rs b/litellm-rust/crates/core/src/messages/client.rs deleted file mode 100644 index ca70b1b03eb..00000000000 --- a/litellm-rust/crates/core/src/messages/client.rs +++ /dev/null @@ -1,14 +0,0 @@ -use std::{sync::OnceLock, time::Duration}; - -use crate::constants::{MESSAGES_CONNECT_TIMEOUT_SECS, MESSAGES_TIMEOUT_SECS}; - -pub(super) fn http_client() -> &'static reqwest::Client { - static CLIENT: OnceLock = OnceLock::new(); - CLIENT.get_or_init(|| { - reqwest::Client::builder() - .timeout(Duration::from_secs(MESSAGES_TIMEOUT_SECS)) - .connect_timeout(Duration::from_secs(MESSAGES_CONNECT_TIMEOUT_SECS)) - .build() - .unwrap_or_else(|_| reqwest::Client::new()) - }) -} diff --git a/litellm-rust/crates/core/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs index d27b79bdc04..fc3bbb36098 100644 --- a/litellm-rust/crates/core/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -1,23 +1,37 @@ use litellm_http::request::string_headers as shared_string_headers; pub(super) use litellm_http::request::truncate_error_body; use litellm_llms::{ - anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG, + anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG, azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG, base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig, + bedrock::messages::invoke_transformations::anthropic_claude3_transformation::BEDROCK_ANTHROPIC_MESSAGES_CONFIG, }; use serde_json::{Map, Value}; +use strum::{EnumString, IntoStaticStr}; use super::Error; const HEADER_CONTEXT: &str = "messages"; -pub(super) fn messages_provider_config( - provider: &str, -) -> Option<&'static dyn BaseAnthropicMessagesConfig> { - match provider { - "anthropic" => Some(&ANTHROPIC_MESSAGES_CONFIG), - "azure_ai" => Some(&AZURE_ANTHROPIC_MESSAGES_CONFIG), - _ => None, +#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)] +#[strum(serialize_all = "snake_case")] +pub(crate) enum MessagesProvider { + Anthropic, + AzureAi, + Bedrock, +} + +impl MessagesProvider { + pub(crate) fn as_str(self) -> &'static str { + self.into() + } + + pub(crate) fn config(self) -> &'static dyn BaseAnthropicMessagesConfig { + match self { + Self::Anthropic => &ANTHROPIC_MESSAGES_CONFIG, + Self::AzureAi => &AZURE_ANTHROPIC_MESSAGES_CONFIG, + Self::Bedrock => &BEDROCK_ANTHROPIC_MESSAGES_CONFIG, + } } } @@ -31,14 +45,26 @@ pub(super) fn string_headers( mod tests { use serde_json::json; - use super::{messages_provider_config, string_headers, truncate_error_body}; + use rstest::rstest; + + use super::{MessagesProvider, string_headers, truncate_error_body}; use crate::messages::Error; + #[rstest] + #[case::anthropic("anthropic", MessagesProvider::Anthropic)] + #[case::azure_ai("azure_ai", MessagesProvider::AzureAi)] + #[case::bedrock("bedrock", MessagesProvider::Bedrock)] + fn provider_round_trips_through_its_python_name( + #[case] name: &str, + #[case] provider: MessagesProvider, + ) { + assert_eq!(name.parse::(), Ok(provider)); + assert_eq!(provider.as_str(), name); + } + #[test] - fn provider_config_resolves_anthropic_and_azure_ai() { - assert!(messages_provider_config("anthropic").is_some()); - assert!(messages_provider_config("azure_ai").is_some()); - assert!(messages_provider_config("openai").is_none()); + fn provider_without_a_messages_config_is_rejected() { + assert!("openai".parse::().is_err()); } #[test] diff --git a/litellm-rust/crates/core/src/messages/error.rs b/litellm-rust/crates/core/src/messages/error.rs deleted file mode 100644 index 2a9723beb38..00000000000 --- a/litellm-rust/crates/core/src/messages/error.rs +++ /dev/null @@ -1,80 +0,0 @@ -use std::sync::Arc; - -use litellm_llms::base_llm::chat::transformation::Error as LlmError; - -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("invalid provider: {0}")] - InvalidProvider(String), - #[error("missing required field: {0}")] - MissingField(&'static str), - #[error("invalid request: {0}")] - InvalidRequest(String), - #[error("invalid response: {0}")] - InvalidResponse(String), - #[error("unsupported by the Rust messages route: {0}")] - Unsupported(&'static str), - #[error(transparent)] - Auth(#[from] litellm_auth::Error), - #[error(transparent)] - Transport(#[from] litellm_http::transport::Error), - #[error(transparent)] - Headers(#[from] litellm_http::request::HeaderError), - #[error(transparent)] - Secret(#[from] SecretError), -} - -#[derive(Clone, Debug, thiserror::Error)] -#[error(transparent)] -pub struct SecretError(Arc); - -impl SecretError { - pub fn source_error(&self) -> &litellm_secrets::Error { - &self.0 - } -} - -impl From for Error { - fn from(error: litellm_secrets::Error) -> Self { - Self::Secret(SecretError(Arc::new(error))) - } -} - -impl PartialEq for SecretError { - fn eq(&self, other: &Self) -> bool { - Arc::ptr_eq(&self.0, &other.0) - } -} - -impl Eq for SecretError {} - -impl From for Error { - fn from(error: LlmError) -> Self { - match error { - error @ LlmError::InvalidType { .. } => Self::InvalidRequest(error.to_string()), - LlmError::MissingField(field) => Self::MissingField(field), - LlmError::InvalidRequest(message) => Self::InvalidRequest(message), - LlmError::InvalidResponse(message) => Self::InvalidResponse(message), - LlmError::Unsupported(reason) => Self::Unsupported(reason), - LlmError::Auth(error) => Self::Auth(error), - } - } -} - -impl Error { - pub fn is_request(&self) -> bool { - match self { - Self::InvalidProvider(_) - | Self::MissingField(_) - | Self::InvalidRequest(_) - | Self::Unsupported(_) - | Self::Headers(_) => true, - Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }), - _ => false, - } - } - - pub fn is_response(&self) -> bool { - matches!(self, Self::InvalidResponse(_)) - } -} diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index de1a5f476ed..6832a59c7bd 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -1,47 +1,143 @@ use std::time::Duration; -use litellm_http::{request::http_request, transport::Error as TransportError}; -use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig; +use bytes::Bytes; +use futures_util::{StreamExt, TryStreamExt, stream::BoxStream}; +use litellm_auth::AuthServices; +use litellm_host::{ + event::{MachineEvent, RawResponse, RequestContext, WireRequest}, + hooks::RouteHooks, +}; +use litellm_http::transport::Error as TransportError; +use litellm_llms::base_llm::{ + anthropic_messages::{ + streaming::{ByteStream, StreamDecoder, encode_anthropic_sse}, + transformation::BaseAnthropicMessagesConfig, + }, + auth::{Authenticated, resolve_auth}, +}; +use litellm_tracing::{ByteChunk, debug}; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use serde_json::Value; -use super::{Error, client::http_client, common_utils::truncate_error_body}; +use super::{ + Error, MessagesResponse, common_utils::truncate_error_body, prepare::ProviderMessagesRequest, +}; +use crate::{constants::MESSAGES_TIMEOUT_SECS, outbound::outbound_request}; -pub(super) fn network(error: reqwest::Error) -> Error { +pub(super) async fn execute( + http: &litellm_http::Client, + auth: &AuthServices, + request: ProviderMessagesRequest, + hooks: &impl RouteHooks, +) -> Result { + let ProviderMessagesRequest { + provider, + url, + body, + environment, + timeout, + api_key, + } = request; + let stream = body.params.stream == Some(true); + let context = RequestContext { + model: body.model.clone(), + custom_llm_provider: provider.as_str().to_string(), + optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?, + secret_fields: Vec::new(), + api_key, + }; + let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?; + let wire = hooks + .before_send( + WireRequest { + url, + headers: authenticated.headers, + body: serde_json::to_value(&body).map_err(serialize_failure)?, + }, + context, + ) + .await?; + let provider_name = provider.as_str(); + debug!(provider = provider_name, stream, body = %wire.body, "provider request"); + let response = send( + http, + Authenticated { + headers: wire.headers, + signer: authenticated.signer, + }, + &wire.url, + &wire.body, + timeout, + ) + .await?; + debug!( + provider = provider_name, + status = response.status().as_u16(), + "provider response headers" + ); + if !response.status().is_success() { + return Err(provider_error(response).await); + } + let config = provider.config(); + if stream { + return Ok(streaming_response( + response, + config.stream_decoder(), + provider_name, + )); + } + let text = response.text().await.map_err(network)?; + debug!(body = text.as_str(), "provider response body"); + hooks + .emit(MachineEvent::ResponseReceived { + raw: RawResponse { body: text.clone() }, + }) + .await?; + decode_response(config, &body.model, &text) + .map(|message| MessagesResponse::Message(Box::new(message))) +} + +fn serialize_failure(err: serde_json::Error) -> Error { + Error::InvalidRequest(format!( + "failed to serialize Anthropic messages request: {err}" + )) +} + +fn network(error: reqwest::Error) -> Error { Error::Transport(TransportError::Network(error.to_string())) } -pub(super) async fn send( +async fn send( + http: &litellm_http::Client, + authenticated: Authenticated, url: &str, - headers: &[(String, String)], body: &Value, timeout: Option, ) -> Result { - let encoded = serde_json::to_vec(body) - .map_err(|err| Error::InvalidRequest(format!("failed to encode messages body: {err}")))?; - let builder = headers.iter().fold( - http_client().post(url).body(encoded), - |builder, (key, value)| builder.header(key, value), - ); - let builder = match timeout { - Some(duration) => builder.timeout(duration), - None => builder, - }; - http_request(builder).await.map_err(network) + let request = outbound_request( + authenticated, + url.to_string(), + body, + Some(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))), + )?; + request.send(http).await.map_err(network) } -pub(super) async fn provider_error(response: reqwest::Response) -> Error { +async fn provider_error(response: reqwest::Response) -> Error { let status = response.status().as_u16(); match response.text().await { - Ok(text) => Error::Transport(TransportError::Http { - status, - body: truncate_error_body(&text), - }), + Ok(text) => { + litellm_tracing::debug!(status, body = text.as_str(), "provider error body"); + Error::Transport(TransportError::Http { + status, + body: truncate_error_body(&text), + }) + } Err(error) => network(error), } } -pub(super) fn decode_response( +fn decode_response( config: &dyn BaseAnthropicMessagesConfig, model: &str, text: &str, @@ -52,3 +148,97 @@ pub(super) fn decode_response( .transform_anthropic_messages_response(model, response) .map_err(Error::from) } + +fn streaming_response( + response: reqwest::Response, + decoder: Option, + provider: &'static str, +) -> MessagesResponse { + let headers = response + .headers() + .iter() + .filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string()))) + .collect(); + let chunks = match decoder { + None => futures_util::stream::try_unfold(response, move |mut response| async move { + let chunk = response.chunk().await.map_err(network)?; + Ok(chunk.map(|chunk| { + log_chunk(provider, "provider_response", &chunk); + (chunk, response) + })) + }) + .boxed(), + Some(decode) => decoded_chunks(response, decode, provider), + }; + MessagesResponse::Stream { headers, chunks } +} + +fn decoded_chunks( + response: reqwest::Response, + decode: StreamDecoder, + provider: &'static str, +) -> BoxStream<'static, Result> { + let bytes: ByteStream = response + .bytes_stream() + .inspect_ok(move |chunk| log_chunk(provider, "provider_response", chunk)) + .map_err(std::io::Error::other) + .boxed(); + futures_util::stream::try_unfold(decode(bytes), move |mut events| async move { + let Some(event) = events.try_next().await? else { + return Ok(None); + }; + let chunk = encode_anthropic_sse(&event)?; + log_chunk(provider, "client_response", &chunk); + Ok(Some((chunk, events))) + }) + .boxed() +} + +fn log_chunk(provider: &str, stage: &str, data: &Bytes) { + let chunk = ByteChunk::new(data); + debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk"); +} + +#[cfg(test)] +mod tests { + use litellm_llms::base_llm::anthropic_messages::streaming::anthropic_sse_event_stream; + use rstest::rstest; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::any}; + + use super::*; + + #[rstest] + #[case::event( + "data: {\"type\":\"ping\"}\n\n", + Some("event: ping\ndata: {\"type\":\"ping\"}\n\n") + )] + #[case::invalid_event("data: invalid\n\ndata: {\"type\":\"ping\"}\n\n", None)] + #[tokio::test] + async fn decoded_streams_encode_events_and_stop_at_the_first_error( + #[case] body: &'static str, + #[case] expected: Option<&str>, + ) { + let upstream = MockServer::start().await; + Mock::given(any()) + .respond_with(ResponseTemplate::new(200).set_body_raw(body, "text/event-stream")) + .mount(&upstream) + .await; + let response = litellm_http::Client::plain_for_test() + .get(upstream.uri()) + .send() + .await + .unwrap(); + let MessagesResponse::Stream { mut chunks, .. } = + streaming_response(response, Some(anthropic_sse_event_stream), "test") + else { + panic!("a streaming response returns chunks"); + }; + + let chunk = chunks.next().await.unwrap(); + match expected { + Some(expected) => assert_eq!(chunk.unwrap().as_ref(), expected.as_bytes()), + None => assert!(matches!(chunk, Err(Error::InvalidResponse(_))), "{chunk:?}"), + } + assert!(chunks.next().await.is_none()); + } +} diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index 180eb08810e..f3a57da4d32 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -1,48 +1,27 @@ -//! The Anthropic Messages call, the Rust equivalent of Python's -//! `litellm.messages()`. +//! The Anthropic Messages call, the Rust equivalent of Python's `litellm.messages()`. //! -//! [`route`] is the call as a machine a host drives, streaming or not. [`messages`] runs -//! it in process for a caller that already holds the request and wants the message. +//! [`messages`] prepares the provider request and sends it in process. [`route`] runs the +//! same two steps as a machine for a host that answers the call's operations itself. -mod error; -pub mod types; -pub use error::Error; -mod client; mod common_utils; mod handler; mod prepare; pub mod route; -use std::sync::Arc; +mod types; -use litellm_secrets::source::EnvironmentSecrets; -use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; -use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine}; -use serde_json::Value; +use litellm_http::{ClientVariant, HttpClientConfig}; +use litellm_secrets::source::SecretSource; -use crate::messages::types::MessagesRequest; +pub use crate::error::RouteError as Error; +pub use types::{MessagesCall, MessagesResponse, MessagesShaping, messages_body}; -pub async fn messages(request: MessagesRequest<'_>) -> Result { - let Value::Object(body) = request.body else { - return Err(Error::InvalidRequest( - "messages body must be an object".into(), - )); - }; - let call = MessagesCall { - model: request.model.into(), - body, - api_key: request.api_key.map(Into::into), - api_base: request.api_base.map(Into::into), - custom_llm_provider: request.custom_llm_provider.map(Into::into), - extra_headers: request.extra_headers, - provider_specific_header: request.provider_specific_header, - timeout: request.timeout, - shaping: request.shaping, - }; - let secrets = Arc::new(EnvironmentSecrets::python_compatible()); - match litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await? { - MessagesOutput::Message(message) => Ok(*message), - MessagesOutput::Streamed => Err(Error::Unsupported( - "streamed responses need a streaming host", - )), - } +pub async fn messages( + resources: &crate::resources::CoreResources, + config: &HttpClientConfig, + secrets: &dyn SecretSource, + call: MessagesCall, +) -> Result { + let http = resources.pool.client(config, ClientVariant::Provider)?; + let request = prepare::prepare(call, secrets).await?; + handler::execute(&http, &resources.auth, request, &()).await } diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index dc4b3562e3f..c8e90cb5f2c 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -1,3 +1,6 @@ +use std::time::Duration; + +use litellm_auth::SecretValue; use litellm_core_utils::{ dot_notation_indexing::delete_nested_value, get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}, @@ -5,30 +8,51 @@ use litellm_core_utils::{ settings::Lookup, }; use litellm_llms::{ - anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request, - base_llm::anthropic_messages::transformation::{ - BaseAnthropicMessagesConfig, MessagesTransformContext, + anthropic::messages::handler::shape_anthropic_messages_request, + base_llm::{ + anthropic_messages::transformation::MessagesTransformContext, + auth::{ValidatedEnvironment, with_default_headers}, }, }; +use litellm_secrets::source::SecretSource; use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; -use serde_json::{Map, Value}; use super::{ - Error, - common_utils::{messages_provider_config, string_headers}, + Error, MessagesCall, + common_utils::{MessagesProvider, string_headers}, + types::invalid_request, }; -use crate::messages::types::{MessagesRequest, ProviderMessagesRequest}; -pub(super) struct ResolvedProvider<'a> { - pub(super) model: &'a str, - pub(super) provider: &'a str, - pub(super) config: &'static dyn BaseAnthropicMessagesConfig, +struct ResolvedProvider { + model: String, + provider: MessagesProvider, } -pub(super) fn resolve_provider<'a>( - model: &'a str, - custom_llm_provider: Option<&'a str>, -) -> Result, Error> { +pub(super) struct ProviderMessagesRequest { + pub(super) provider: MessagesProvider, + pub(super) url: String, + pub(super) body: AnthropicMessagesRequest, + pub(super) environment: ValidatedEnvironment, + pub(super) timeout: Option, + /// The caller's own credential, reported to the host beside the wire request. + pub(super) api_key: Option, +} + +pub(super) async fn prepare( + call: MessagesCall, + secrets: &dyn SecretSource, +) -> Result { + let resolved = resolve_provider(&call.body.model, call.custom_llm_provider.as_deref())?; + let secrets = secrets + .resolve(resolved.provider.config().secret_names()) + .await?; + prepare_provider_request(call, resolved, secrets.as_ref()) +} + +fn resolve_provider( + model: &str, + custom_llm_provider: Option<&str>, +) -> Result { let CustomLlmProvider { model, custom_llm_provider: provider, @@ -44,82 +68,79 @@ pub(super) fn resolve_provider<'a>( "unable to resolve custom_llm_provider for messages request".to_string(), ) })?; - let config = messages_provider_config(provider) - .ok_or_else(|| Error::InvalidProvider(provider.to_string()))?; + let provider = provider + .parse() + .map_err(|_| Error::InvalidProvider(provider.to_string()))?; Ok(ResolvedProvider { - model, + model: model.to_string(), provider, - config, }) } -pub(super) fn prepare_provider_request( - request: MessagesRequest<'_>, - resolved: ResolvedProvider<'_>, +fn prepare_provider_request( + call: MessagesCall, + resolved: ResolvedProvider, secrets: &dyn Lookup, ) -> Result { - let ResolvedProvider { - model, - provider, - config, - } = resolved; - let model = model.to_string(); + let ResolvedProvider { model, provider } = resolved; + let MessagesCall { + body, + api_key, + api_base, + extra_headers, + provider_specific_header, + timeout, + shaping, + .. + } = call; + let config = provider.config(); let env_lookup = |key: &str| secrets.get(key); - let typed_request: AnthropicMessagesRequest = - serde_json::from_value(request.body).map_err(invalid_request)?; let sanitized = shape_anthropic_messages_request( - AnthropicMessagesRequest { - model: model.clone(), - ..typed_request - }, - request.shaping.reasoning_auto_summary, + AnthropicMessagesRequest { model, ..body }, + shaping.reasoning_auto_summary, )?; - let trimmed = - without_additional_drop_params(sanitized, &request.shaping.additional_drop_params)?; + let trimmed = without_additional_drop_params(sanitized, &shaping.additional_drop_params)?; let transformed = config.transform_anthropic_messages_request( trimmed, - &MessagesTransformContext::new(request.shaping.capabilities, request.shaping.drop_params), + &MessagesTransformContext::new(shaping.capabilities, shaping.drop_params), )?; - let scoped = get_provider_specific_headers(request.provider_specific_header.as_ref(), provider); + let scoped = + get_provider_specific_headers(provider_specific_header.as_ref(), provider.as_str()); let forwarded = string_headers(Some( - request - .extra_headers - .into_iter() - .flatten() - .chain(scoped) - .collect(), + extra_headers.into_iter().flatten().chain(scoped).collect(), ))?; - let authenticated = config.authenticate(forwarded, request.api_key, &env_lookup)?; - let headers = config.request_headers( - with_default_headers(authenticated, config.default_headers()), - &transformed, - ); + let validated = config.validate_environment( + forwarded, + api_key.as_deref(), + &transformed.model, + &env_lookup, + )?; + let environment = ValidatedEnvironment { + headers: config.request_headers( + with_default_headers(validated.headers, config.default_headers()), + &transformed, + ), + auth: validated.auth, + }; - let body = serde_json::to_value(transformed).map_err(|err| { - Error::InvalidRequest(format!( - "failed to serialize Anthropic messages request: {err}" - )) - })?; - - let url = config.get_complete_url(request.api_base, &model, &env_lookup)?; + let url = if transformed.params.stream == Some(true) { + config.complete_stream_url(api_base.as_deref(), &transformed.model, &env_lookup)? + } else { + config.get_complete_url(api_base.as_deref(), &transformed.model, &env_lookup)? + }; Ok(ProviderMessagesRequest { - provider: provider.to_string(), - model, - config, + provider, url, - body, - upstream_headers: headers, - timeout: request.timeout, + body: transformed, + environment, + timeout, + api_key: api_key.map(SecretValue::new), }) } -fn invalid_request(err: serde_json::Error) -> Error { - Error::InvalidRequest(format!("invalid Anthropic messages request: {err}")) -} - fn without_additional_drop_params( request: AnthropicMessagesRequest, paths: &[String], @@ -127,64 +148,59 @@ fn without_additional_drop_params( if paths.is_empty() { return Ok(request); } - let Value::Object(fields) = serde_json::to_value(request).map_err(invalid_request)? else { - return Err(Error::InvalidRequest( - "Anthropic messages request did not serialize to an object".to_string(), - )); - }; - let (required, optional): (Map, Map) = fields - .into_iter() - .partition(|(key, _)| matches!(key.as_str(), "model" | "messages")); - let trimmed = paths.iter().fold(Value::Object(optional), |body, path| { - delete_nested_value(body, path) - }); - let merged: Map = required - .into_iter() - .chain(trimmed.as_object().cloned().unwrap_or_default()) - .collect(); - serde_json::from_value(Value::Object(merged)).map_err(invalid_request) -} - -fn with_default_headers( - headers: Vec<(String, String)>, - defaults: &[(&str, &str)], -) -> Vec<(String, String)> { - let missing: Vec<(String, String)> = defaults + let params = serde_json::to_value(request.params).map_err(invalid_request)?; + let trimmed = paths .iter() - .filter(|(name, _)| { - !headers - .iter() - .any(|(header, _)| header.eq_ignore_ascii_case(name)) - }) - .map(|(name, value)| ((*name).to_string(), (*value).to_string())) - .collect(); - headers.into_iter().chain(missing).collect() + .fold(params, |params, path| delete_nested_value(params, path)); + Ok(AnthropicMessagesRequest { + params: serde_json::from_value(trimmed).map_err(invalid_request)?, + ..request + }) } #[cfg(test)] mod tests { + use litellm_llms::base_llm::auth::resolve_auth; use litellm_types::utils::ProviderSpecificHeaders; use rstest::{fixture, rstest}; - use serde_json::json; + use serde_json::{Map, Value, json}; use super::*; - use crate::messages::types::MessagesShaping; + use crate::messages::MessagesShaping; #[fixture] fn shaping() -> MessagesShaping { MessagesShaping::default() } - fn prepare(request: MessagesRequest<'_>) -> Result { - prepare_with_secrets(request, &|_: &str| None) + fn body(value: Value) -> AnthropicMessagesRequest { + serde_json::from_value(value).unwrap() + } + + fn prepare(call: MessagesCall) -> Result { + prepare_with_secrets(call, &|_: &str| None) } fn prepare_with_secrets( - request: MessagesRequest<'_>, + call: MessagesCall, secrets: &dyn Lookup, ) -> Result { - let resolved = resolve_provider(request.model, request.custom_llm_provider)?; - prepare_provider_request(request, resolved, secrets) + let resolved = resolve_provider(&call.body.model, call.custom_llm_provider.as_deref())?; + prepare_provider_request(call, resolved, secrets) + } + + /// The headers as they go on the wire, credential applied. + fn wire_headers(prepared: &ProviderMessagesRequest) -> Vec<(String, String)> { + tokio::runtime::Builder::new_current_thread() + .build() + .unwrap() + .block_on(resolve_auth( + &litellm_auth::AuthServices::default(), + prepared.environment.clone(), + &|_| None, + )) + .unwrap() + .headers } #[rstest] @@ -221,12 +237,13 @@ mod tests { .map(|(_, value)| value.to_string()) }; let prepared = prepare_with_secrets( - MessagesRequest { - model: "claude-test", - body: json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}), + MessagesCall { + body: body( + json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}), + ), api_key: None, api_base: None, - custom_llm_provider: Some("anthropic"), + custom_llm_provider: Some("anthropic".into()), extra_headers: None, provider_specific_header: None, timeout: None, @@ -235,8 +252,8 @@ mod tests { &lookup, ) .unwrap(); - let auth: Vec<(&str, &str)> = prepared - .upstream_headers + let headers = wire_headers(&prepared); + let auth: Vec<(&str, &str)> = headers .iter() .filter(|(name, _)| matches!(name.as_str(), "x-api-key" | "authorization")) .map(|(name, value)| (name.as_str(), value.as_str())) @@ -247,48 +264,18 @@ mod tests { ); } - fn prepared_body(body: Value, shaping: MessagesShaping) -> Result { - prepare(MessagesRequest { - model: "anthropic/claude-test", - body, - api_key: Some("sk-test"), - api_base: Some("https://anthropic.test"), - custom_llm_provider: Some("anthropic"), + fn prepared_body(fields: Value, shaping: MessagesShaping) -> Result { + prepare(MessagesCall { + body: body(fields), + api_key: Some("sk-test".into()), + api_base: Some("https://anthropic.test".into()), + custom_llm_provider: Some("anthropic".into()), extra_headers: None, provider_specific_header: None, timeout: None, shaping, }) - .map(|prepared| prepared.body) - } - - #[rstest] - #[case::nothing_forwarded( - &[], - &[("x-version", "1"), ("content-type", "application/json")], - &[("x-version", "1"), ("content-type", "application/json")], - )] - #[case::forwarded_header_wins_in_any_case( - &[("X-Version", "custom"), ("x-api-key", "k")], - &[("x-version", "1"), ("content-type", "application/json")], - &[("X-Version", "custom"), ("x-api-key", "k"), ("content-type", "application/json")], - )] - #[case::no_defaults(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])] - fn default_headers_fill_only_missing_names( - #[case] forwarded: &[(&str, &str)], - #[case] defaults: &[(&str, &str)], - #[case] expected: &[(&str, &str)], - ) { - let owned = |headers: &[(&str, &str)]| -> Vec<(String, String)> { - headers - .iter() - .map(|(name, value)| ((*name).to_string(), (*value).to_string())) - .collect() - }; - assert_eq!( - with_default_headers(owned(forwarded), defaults), - owned(expected) - ); + .map(|prepared| serde_json::to_value(prepared.body).unwrap()) } #[rstest] @@ -380,20 +367,22 @@ mod tests { {"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "anthropic", "x-priority": "scoped"}} ])) .unwrap(); - let prepared = prepare(MessagesRequest { - model, - body: json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}), - api_key: Some("sk-test"), - api_base: Some("https://resource.services.ai.azure.com"), - custom_llm_provider, - extra_headers: Some(serde_json::from_value(json!({"x-priority": "extra"})).unwrap()), + let prepared = prepare(MessagesCall { + body: body( + json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}), + ), + api_key: Some("sk-test".into()), + api_base: Some("https://resource.services.ai.azure.com".into()), + custom_llm_provider: custom_llm_provider.map(Into::into), + extra_headers: Some(Map::from_iter([("x-priority".into(), json!("extra"))])), provider_specific_header: Some(configured), timeout: None, shaping, }) .unwrap(); let caller_headers: Vec<(&str, &str)> = prepared - .upstream_headers + .environment + .headers .iter() .filter(|(name, _)| matches!(name.as_str(), "x-priority" | "x-scoped")) .map(|(name, value)| (name.as_str(), value.as_str())) diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 40aff185e81..5e5bbc927b6 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -1,50 +1,20 @@ use std::{ convert::Infallible, sync::{Arc, Mutex}, - time::Duration, }; use bytes::Bytes; -use litellm_auth::SecretValue; +use futures_util::TryStreamExt; use litellm_host::{ - event::{MachineEvent, RawResponse, RequestContext, WireRequest}, host::{Demand, Host}, machine::{CallMachine, HostChannel, MachineFault}, protocol::Protocol, }; +use litellm_http::{Client, ClientVariant, HttpClientConfig}; use litellm_secrets::source::SecretSource; -use litellm_types::{ - llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse, - utils::ProviderSpecificHeaders, -}; -use serde_json::{Map, Value}; +use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; -use super::{ - Error, - handler::{decode_response, network, provider_error, send}, - prepare::{prepare_provider_request, resolve_provider}, - types::{MessagesRequest, MessagesShaping}, -}; -use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; - -/// The caller's request as the host projects it. -pub struct MessagesCall { - pub model: String, - pub body: Map, - pub api_key: Option, - pub api_base: Option, - pub custom_llm_provider: Option, - pub extra_headers: Option>, - pub provider_specific_header: Option, - pub timeout: Option, - pub shaping: MessagesShaping, -} - -impl MessagesCall { - fn streams(&self) -> bool { - self.body.get("stream").and_then(Value::as_bool) == Some(true) - } -} +use super::{Error, MessagesCall, MessagesResponse, handler::execute, prepare::prepare}; pub enum MessagesOutput { Message(Box), @@ -108,98 +78,43 @@ impl Host for LocalMessagesHost { } } -pub fn messages_machine(secrets: Arc) -> MessagesMachine { - CallMachine::new(move |host| Box::pin(execute(host, secrets.clone()))) +pub fn messages_machine( + resources: &crate::resources::CoreResources, + config: &HttpClientConfig, + secrets: Arc, +) -> Result { + let http = resources.pool.client(config, ClientVariant::Provider)?; + let auth = resources.auth.clone(); + Ok(CallMachine::new(move |host| { + Box::pin(drive(host, http, auth, secrets)) + })) } -async fn execute( +/// The call as its host sees it: projection first, then the same prepare and execute as +/// [`super::messages`], with each chunk of a stream handed over as it arrives. +async fn drive( host: MessagesHost, + http: Client, + auth: Arc, secrets: Arc, ) -> Result { let call = host.project().await?; - let stream = call.streams(); - let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?; - let secrets = secrets.resolve(resolved.config.secret_names()).await?; - let request = prepare_provider_request( - MessagesRequest { - model: &call.model, - body: Value::Object(call.body.clone()), - api_key: call.api_key.as_deref(), - api_base: call.api_base.as_deref(), - custom_llm_provider: call.custom_llm_provider.as_deref(), - extra_headers: call.extra_headers.clone(), - provider_specific_header: call.provider_specific_header.clone(), - timeout: call.timeout, - shaping: call.shaping.clone(), - }, - resolved, - secrets.as_ref(), - )?; - if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER { - return Err(Error::Unsupported("streaming messages for this provider")); - } - let context = RequestContext { - model: request.model.clone(), - custom_llm_provider: request.provider.clone(), - optional_params: Value::Object( - request - .body - .as_object() - .into_iter() - .flatten() - .filter(|(name, _)| !matches!(name.as_str(), "model" | "messages")) - .map(|(name, value)| (name.clone(), value.clone())) - .collect(), - ), - secret_fields: Vec::new(), - api_key: call.api_key.clone().map(SecretValue::new), - }; - let wire = host - .before_send( - WireRequest { - url: request.url, - headers: request.upstream_headers, - body: request.body, - }, - context, - ) - .await?; - let response = send(&wire.url, &wire.headers, &wire.body, request.timeout).await?; - if !response.status().is_success() { - return Err(provider_error(response).await); - } - if stream { - return relay(&host, response).await; - } - let text = response.text().await.map_err(network)?; - host.emit(MachineEvent::ResponseReceived { - raw: RawResponse { body: text.clone() }, - }) - .await?; - decode_response(request.config, &request.model, &text) - .map(|message| MessagesOutput::Message(Box::new(message))) -} - -/// Hands each upstream chunk to the caller as it arrives. A caller that stops reading -/// ends the upstream read, and the call completes with what it delivered. -async fn relay( - host: &MessagesHost, - mut response: reqwest::Response, -) -> Result { - let head = MessagesStreamHead { - headers: response - .headers() - .iter() - .filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string()))) - .collect(), - }; - if host.open(head).await? == Demand::Detached { - return Ok(MessagesOutput::Streamed); - } - while let Some(chunk) = response.chunk().await.map_err(network)? { - if host.deliver(chunk).await? == Demand::Detached { - break; + let request = prepare(call, secrets.as_ref()).await?; + match execute(&http, &auth, request, &host).await? { + MessagesResponse::Message(message) => Ok(MessagesOutput::Message(message)), + MessagesResponse::Stream { + headers, + mut chunks, + } => { + if host.open(MessagesStreamHead { headers }).await? == Demand::Detached { + return Ok(MessagesOutput::Streamed); + } + while let Some(chunk) = chunks.try_next().await? { + if host.deliver(chunk).await? == Demand::Detached { + break; + } + } + Ok(MessagesOutput::Streamed) } } - Ok(MessagesOutput::Streamed) } diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index 4a5dd2926e0..b09cb96a919 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -1,13 +1,46 @@ use std::time::Duration; -use litellm_llms::{ - anthropic::common_utils::AnthropicModelCapabilities, - base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig, +use bytes::Bytes; +use futures_util::stream::BoxStream; +use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities; +use litellm_types::{ + llms::anthropic_messages::{ + anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse, + }, + utils::ProviderSpecificHeaders, }; -use litellm_types::utils::ProviderSpecificHeaders; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use super::Error; + +pub struct MessagesCall { + pub body: AnthropicMessagesRequest, + pub api_key: Option, + pub api_base: Option, + pub custom_llm_provider: Option, + pub extra_headers: Option>, + pub provider_specific_header: Option, + pub timeout: Option, + pub shaping: MessagesShaping, +} + +pub fn messages_body(body: Map) -> Result { + serde_json::from_value(Value::Object(body)).map_err(invalid_request) +} + +pub(super) fn invalid_request(err: serde_json::Error) -> Error { + Error::InvalidRequest(format!("invalid Anthropic messages request: {err}")) +} + +pub enum MessagesResponse { + Message(Box), + Stream { + headers: Vec<(String, String)>, + chunks: BoxStream<'static, Result>, + }, +} + #[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] pub struct MessagesShaping { #[serde(default)] @@ -20,33 +53,11 @@ pub struct MessagesShaping { pub additional_drop_params: Vec, } -pub struct MessagesRequest<'a> { - pub model: &'a str, - pub body: Value, - pub api_key: Option<&'a str>, - pub api_base: Option<&'a str>, - pub custom_llm_provider: Option<&'a str>, - pub extra_headers: Option>, - pub provider_specific_header: Option, - pub timeout: Option, - pub shaping: MessagesShaping, -} - -pub struct ProviderMessagesRequest { - pub provider: String, - pub model: String, - pub config: &'static dyn BaseAnthropicMessagesConfig, - pub url: String, - pub body: Value, - pub upstream_headers: Vec<(String, String)>, - pub timeout: Option, -} - #[cfg(test)] mod tests { use litellm_llms::anthropic::common_utils::SupportedEffortTiers; use rstest::rstest; - use serde_json::json; + use serde_json::{Value, json}; use super::*; diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 18961ec96fa..f13d6984763 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -110,7 +110,10 @@ mod tests { } fn client() -> OcrClient { - OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new()) + OcrClient::for_test( + litellm_http::Client::plain_for_test(), + litellm_http::Client::no_redirect_for_test(), + ) } fn request(model: &str, base: &str, document: Value, options: Value) -> LiteLLMOcrRequest { diff --git a/litellm-rust/crates/core/src/outbound.rs b/litellm-rust/crates/core/src/outbound.rs index 7fc90084e6f..0cdbb465f60 100644 --- a/litellm-rust/crates/core/src/outbound.rs +++ b/litellm-rust/crates/core/src/outbound.rs @@ -1,30 +1,20 @@ use std::time::Duration; -use litellm_auth::RequestAuth; -use litellm_auth_aws::SigV4Signer; use litellm_http::outbound::OutboundRequest; -use serde_json::{Map, Value}; +use litellm_llms::base_llm::auth::Authenticated; +use serde_json::Value; /// Header credentials are already in `headers`; SigV4 is applied here, over the /// bytes that are sent. -pub(crate) async fn outbound_request( - auth: &RequestAuth, +pub(crate) fn outbound_request( + authenticated: Authenticated, url: String, - headers: Vec<(String, String)>, body: &Value, timeout: Option, - optional_params: &Map, -) -> Result -where - E: From + From, -{ - let RequestAuth::AwsSigV4 { region, service } = auth else { - return Ok(OutboundRequest::json(url, headers, body, timeout)?); - }; - let env_lookup = |key: &str| std::env::var(key).ok(); - let signer = - SigV4Signer::resolve(region.clone(), service, optional_params, &env_lookup).await?; - Ok(OutboundRequest::signed_json( - url, headers, body, timeout, &signer, - )?) +) -> Result { + let Authenticated { headers, signer } = authenticated; + match signer { + None => OutboundRequest::json(url, headers, body, timeout), + Some(signer) => OutboundRequest::signed_json(url, headers, body, timeout, &signer), + } } diff --git a/litellm-rust/crates/core/src/resources.rs b/litellm-rust/crates/core/src/resources.rs new file mode 100644 index 00000000000..37a29502649 --- /dev/null +++ b/litellm-rust/crates/core/src/resources.rs @@ -0,0 +1,38 @@ +use std::sync::Arc; + +use litellm_auth::AuthServices; +use litellm_http::{HttpClientConfig, HttpClientPool, media::UrlPolicy}; +use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings}; +use litellm_secrets::source::SecretSource; + +#[derive(Clone)] +pub struct CoreResources { + pub pool: Arc, + pub auth: Arc, +} + +impl CoreResources { + pub fn new(pool: Arc) -> Self { + Self { + pool, + auth: Arc::new(AuthServices::default()), + } + } + + pub fn ocr_client( + &self, + config: &HttpClientConfig, + url_policy: UrlPolicy, + settings: OcrSettings, + secrets: Arc, + ) -> Result { + OcrClient::new( + &self.pool, + config, + url_policy, + self.auth.clone(), + settings, + secrets, + ) + } +} diff --git a/litellm-rust/crates/core/src/responses/error.rs b/litellm-rust/crates/core/src/responses/error.rs deleted file mode 100644 index 1c940d8ed9b..00000000000 --- a/litellm-rust/crates/core/src/responses/error.rs +++ /dev/null @@ -1,17 +0,0 @@ -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("invalid provider: {0}")] - InvalidProvider(String), - #[error("invalid request: {0}")] - InvalidRequest(String), - #[error("invalid response: {0}")] - InvalidResponse(String), - #[error("routing error: {0}")] - Routing(String), - #[error(transparent)] - Auth(#[from] litellm_auth::Error), - #[error(transparent)] - Transport(#[from] litellm_http::transport::Error), - #[error(transparent)] - Headers(#[from] litellm_http::request::HeaderError), -} diff --git a/litellm-rust/crates/core/src/responses/mod.rs b/litellm-rust/crates/core/src/responses/mod.rs index bc0f71896e5..464a81fe89c 100644 --- a/litellm-rust/crates/core/src/responses/mod.rs +++ b/litellm-rust/crates/core/src/responses/mod.rs @@ -1,3 +1,2 @@ -mod error; -pub use error::Error; +pub use crate::error::RouteError as Error; pub mod websocket; diff --git a/litellm-rust/crates/core/tests/audio_transcription.rs b/litellm-rust/crates/core/tests/audio_transcription.rs index 196f085a6c3..c4dfea87319 100644 --- a/litellm-rust/crates/core/tests/audio_transcription.rs +++ b/litellm-rust/crates/core/tests/audio_transcription.rs @@ -10,6 +10,10 @@ use support::*; const MODEL: &str = "mistral.voxtral-mini-3b-2507"; +async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result { + audio_transcription(&support::resources(), &http_config(), request).await +} + fn transcript_response(text: &str) -> ResponseTemplate { json_response(json!({"output": {"message": {"content": [{"text": text}]}}})) } @@ -47,7 +51,7 @@ async fn bedrock_converse_request_is_signed_for_the_requested_region( let upstream = upstream([transcript_response("hello")]).await; let base = upstream.uri(); - let response = audio_transcription(AudioTranscriptionRequest { + let response = transcribe(AudioTranscriptionRequest { api_base: Some(&base), optional_params: aws_params(region), ..request @@ -79,7 +83,7 @@ async fn the_provider_can_come_from_the_model_prefix(request: AudioTranscription let base = upstream.uri(); let model = format!("bedrock/{MODEL}"); - audio_transcription(AudioTranscriptionRequest { + transcribe(AudioTranscriptionRequest { model: &model, custom_llm_provider: None, api_base: Some(&base), @@ -110,7 +114,7 @@ async fn audio_and_transcription_params_reach_the_converse_body( ]) .collect(); - audio_transcription(AudioTranscriptionRequest { + transcribe(AudioTranscriptionRequest { audio: json!({"data": "AQI=", "format": format}), api_base: Some(&base), optional_params, @@ -142,7 +146,7 @@ async fn invalid_audio_is_rejected_before_sending( let upstream = upstream([transcript_response("hello")]).await; let base = upstream.uri(); - let error = audio_transcription(AudioTranscriptionRequest { + let error = transcribe(AudioTranscriptionRequest { audio, api_base: Some(&base), ..request @@ -174,7 +178,7 @@ async fn unsupported_providers_are_rejected_before_sending( #[case] provider: Option<&'static str>, #[case] reported: &str, ) { - let error = audio_transcription(AudioTranscriptionRequest { + let error = transcribe(AudioTranscriptionRequest { model, custom_llm_provider: provider, api_base: Some(UNREACHABLE_BASE), @@ -189,7 +193,7 @@ async fn unsupported_providers_are_rejected_before_sending( #[rstest] #[tokio::test] async fn a_non_string_extra_header_is_rejected(request: AudioTranscriptionRequest<'static>) { - let error = audio_transcription(AudioTranscriptionRequest { + let error = transcribe(AudioTranscriptionRequest { extra_headers: Some(Map::from_iter([("x-count".to_string(), json!(3))])), api_base: Some(UNREACHABLE_BASE), ..request @@ -212,7 +216,7 @@ async fn an_upstream_error_keeps_its_status_and_body( upstream([ResponseTemplate::new(status).set_body_string("upstream said no")]).await; let base = upstream.uri(); - let error = audio_transcription(AudioTranscriptionRequest { + let error = transcribe(AudioTranscriptionRequest { api_base: Some(&base), ..request }) @@ -239,7 +243,7 @@ async fn an_unreadable_success_body_is_an_invalid_response( let upstream = upstream([response]).await; let base = upstream.uri(); - let error = audio_transcription(AudioTranscriptionRequest { + let error = transcribe(AudioTranscriptionRequest { api_base: Some(&base), ..request }) diff --git a/litellm-rust/crates/core/tests/chat_completions.rs b/litellm-rust/crates/core/tests/chat_completions.rs index ae96509fe2e..f5802f8e305 100644 --- a/litellm-rust/crates/core/tests/chat_completions.rs +++ b/litellm-rust/crates/core/tests/chat_completions.rs @@ -4,6 +4,7 @@ use litellm_core::chat_completions::{ Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest, }; use litellm_http::transport::Error as TransportError; +use litellm_types::utils::ChatCompletionsResponse; use rstest::{fixture, rstest}; use serde_json::{Map, Value, json}; use wiremock::ResponseTemplate; @@ -13,6 +14,10 @@ use support::*; const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#; +async fn complete(request: ChatCompletionsRequest<'_>) -> Result { + chat_completions(&support::resources(), &http_config(), request).await +} + fn object(value: Value) -> Map { let Value::Object(map) = value else { panic!("expected a json object, got {value}"); @@ -50,7 +55,7 @@ async fn anthropic_round_trip_translates_the_conversation_and_normalizes_the_res let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; let base = upstream.uri(); - let response = chat_completions(ChatCompletionsRequest { + let response = complete(ChatCompletionsRequest { messages: json!([ {"role": "system", "content": "be terse"}, {"role": "user", "content": "hi"} @@ -90,7 +95,7 @@ async fn the_deployment_key_replaces_a_caller_supplied_x_api_key( let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; let base = upstream.uri(); - chat_completions(ChatCompletionsRequest { + complete(ChatCompletionsRequest { api_base: Some(&base), extra_headers: Some(object( json!({"x-api-key": "caller-key", "x-trace": "kept"}), @@ -116,7 +121,7 @@ async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsReq .await; let base = upstream.uri(); - let response = chat_completions(ChatCompletionsRequest { + let response = complete(ChatCompletionsRequest { model: "bedrock/anthropic.claude-sonnet-4-5", optional_params: object(json!({ "aws_access_key_id": "access-key", @@ -167,7 +172,7 @@ async fn a_response_it_cannot_normalize_is_reported_as_already_sent( let upstream = upstream([anthropic_response(body)]).await; let base = upstream.uri(); - let error = chat_completions(ChatCompletionsRequest { + let error = complete(ChatCompletionsRequest { api_base: Some(&base), ..request }) @@ -188,7 +193,7 @@ async fn an_upstream_error_status_keeps_its_code_and_body( let upstream = upstream([ResponseTemplate::new(status).set_body_string("slow down")]).await; let base = upstream.uri(); - let error = chat_completions(ChatCompletionsRequest { + let error = complete(ChatCompletionsRequest { api_base: Some(&base), ..request }) @@ -210,7 +215,7 @@ async fn an_upstream_error_status_keeps_its_code_and_body( async fn a_connection_that_is_never_established_declines_instead_of_failing( request: ChatCompletionsRequest<'static>, ) { - let error = chat_completions(ChatCompletionsRequest { + let error = complete(ChatCompletionsRequest { api_base: Some(UNREACHABLE_BASE), ..request }) @@ -232,7 +237,7 @@ async fn a_timeout_after_sending_is_not_a_pre_send_decline( upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await; let base = upstream.uri(); - let error = chat_completions(ChatCompletionsRequest { + let error = complete(ChatCompletionsRequest { api_base: Some(&base), timeout: Some(Duration::from_millis(100)), ..request @@ -307,7 +312,7 @@ async fn a_declined_request_fails_the_call_before_sending( let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; let base = upstream.uri(); - let error = chat_completions(ChatCompletionsRequest { + let error = complete(ChatCompletionsRequest { optional_params: object(json!({"stream": true})), api_base: Some(&base), ..request diff --git a/litellm-rust/crates/core/tests/messages/host.rs b/litellm-rust/crates/core/tests/messages/host.rs index ca2aece5ebd..b19ecf11f09 100644 --- a/litellm-rust/crates/core/tests/messages/host.rs +++ b/litellm-rust/crates/core/tests/messages/host.rs @@ -78,7 +78,7 @@ impl Host for RecordingHost { } async fn run_through(host: &RecordingHost) -> Result { - litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await + litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await } fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall { @@ -161,10 +161,10 @@ async fn no_raw_response_is_emitted_for_a_stream_or_a_failure( #[case] response: ResponseTemplate, ) { let upstream = upstream([response]).await; - let mut body = call.body.clone(); - body.insert("stream".into(), json!(true)); - let host = - RecordingHost::passthrough(authenticated(MessagesCall { body, ..call }, upstream.uri())); + let host = RecordingHost::passthrough(authenticated( + with_fields(call, json!({"stream": true})), + upstream.uri(), + )); let _ = run_through(&host).await; @@ -180,15 +180,8 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages call: MessagesCall, ) { let upstream = upstream([message_response()]).await; - let body: Map = call - .body - .clone() - .into_iter() - .chain([("temperature".to_string(), json!(0.2))]) - .collect(); let host = RecordingHost::passthrough(authenticated( MessagesCall { - body, shaping: MessagesShaping { capabilities: AnthropicModelCapabilities { supports_sampling_params: false, @@ -197,7 +190,7 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages drop_params: true, ..MessagesShaping::default() }, - ..call + ..with_fields(call, json!({"temperature": 0.2})) }, upstream.uri(), )); diff --git a/litellm-rust/crates/core/tests/messages/main.rs b/litellm-rust/crates/core/tests/messages/main.rs index 21ee678ced3..534af6d7d06 100644 --- a/litellm-rust/crates/core/tests/messages/main.rs +++ b/litellm-rust/crates/core/tests/messages/main.rs @@ -1,11 +1,14 @@ use std::{sync::Arc, time::Duration}; use litellm_core::messages::{ - Error, - route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine}, - types::MessagesShaping, + Error, MessagesCall, MessagesShaping, + route::{LocalMessagesHost, MessagesMachine, MessagesOutput, messages_machine}, +}; +use litellm_http::{HttpSettings, Resolution}; +use litellm_secrets::source::SecretSource; +use litellm_types::llms::anthropic_messages::{ + anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse, }; -use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use rstest::fixture; use serde_json::{Map, Value, json}; use wiremock::ResponseTemplate; @@ -29,6 +32,24 @@ fn object(value: Value) -> Map { map } +fn body(value: Value) -> AnthropicMessagesRequest { + serde_json::from_value(value).unwrap() +} + +fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall { + let current = object(serde_json::to_value(&call.body).unwrap()); + MessagesCall { + body: body(Value::Object( + current.into_iter().chain(object(fields)).collect(), + )), + ..call + } +} + +fn with_model(call: MessagesCall, model: &str) -> MessagesCall { + with_fields(call, json!({"model": model})) +} + fn message_body() -> Value { json!({ "id": "msg_1", @@ -50,8 +71,7 @@ fn message_response() -> ResponseTemplate { #[fixture] fn call() -> MessagesCall { MessagesCall { - model: MODEL.into(), - body: object(json!({ + body: body(json!({ "model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}] @@ -75,11 +95,16 @@ fn headers<'a>(pairs: impl IntoIterator) -> Option) -> MessagesMachine { + messages_machine(&support::resources(), &http_config(), secrets) + .expect("default HTTP settings build a client") +} + async fn run_with( secrets: Arc, call: MessagesCall, ) -> Result { - litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await + litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await } /// Runs the route with a secret source that knows nothing, so no environment leaks in. diff --git a/litellm-rust/crates/core/tests/messages/request.rs b/litellm-rust/crates/core/tests/messages/request.rs index d37910d4ac4..f6ee0e6dfbf 100644 --- a/litellm-rust/crates/core/tests/messages/request.rs +++ b/litellm-rust/crates/core/tests/messages/request.rs @@ -1,7 +1,5 @@ -use litellm_llms::anthropic::common_utils::{ - ANTHROPIC_ADVISOR_TOOL_TYPE, ANTHROPIC_OAUTH_BETA_HEADER, AnthropicModelCapabilities, - SupportedEffortTiers, beta, -}; +use litellm_llms::anthropic::common_utils::{AnthropicModelCapabilities, SupportedEffortTiers}; +use litellm_types::llms::anthropic::{AnthropicBeta, BetaSet}; use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders}; use rstest::rstest; @@ -124,11 +122,10 @@ async fn each_provider_posts_to_its_messages_endpoint( let upstream = upstream([message_response()]).await; run_message(MessagesCall { - model: model.into(), custom_llm_provider: provider.map(Into::into), api_key: Some("sk".into()), api_base: Some(format!("{}{base_suffix}", upstream.uri())), - ..call + ..with_model(call, model) }) .await; @@ -155,11 +152,10 @@ async fn unsupported_providers_are_rejected_before_sending( #[case] reported: &str, ) { let error = run(MessagesCall { - model: model.into(), custom_llm_provider: provider.map(Into::into), api_key: Some("sk".into()), api_base: Some(UNREACHABLE_BASE.into()), - ..call + ..with_model(call, model) }) .await .err() @@ -206,7 +202,7 @@ async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCa custom_llm_provider: Some("azure_ai".into()), api_key: Some("sk-azure".into()), api_base: Some(upstream.uri()), - body: object(json!({ + body: body(json!({ "model": MODEL, "max_tokens": 16, "messages": [{ @@ -232,19 +228,15 @@ async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCa #[tokio::test] async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall) { let upstream = upstream([message_response()]).await; - let mut body = call.body.clone(); - body.insert("temperature".into(), json!(0.5)); - body.insert("top_k".into(), json!(3)); run_message(MessagesCall { api_key: Some("sk".into()), api_base: Some(upstream.uri()), - body, shaping: MessagesShaping { additional_drop_params: vec!["temperature".into()], ..MessagesShaping::default() }, - ..call + ..with_fields(call, json!({"temperature": 0.5, "top_k": 3})) }) .await; @@ -253,46 +245,37 @@ async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall) assert_eq!(sent["top_k"], 3); } -fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall { - let body: Map = call.body.into_iter().chain(object(fields)).collect(); - MessagesCall { body, ..call } -} - -fn sent_betas(request: &wiremock::Request) -> Vec { +fn sent_betas(request: &wiremock::Request) -> BetaSet { let [header] = <[&str; 1]>::try_from(request.header_values("anthropic-beta")) .unwrap_or_else(|values| panic!("expected one anthropic-beta header, got {values:?}")); - header - .split(',') - .map(str::trim) - .map(str::to_string) - .collect() + header.parse().unwrap() } #[rstest] -#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])] -#[case::fast_mode(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])] -#[case::compaction(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])] +#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[AnthropicBeta::StructuredOutputs20251113])] +#[case::fast_mode(json!({"speed": "fast"}), &[AnthropicBeta::FastMode20260201])] +#[case::compaction(json!({"compaction": {"enabled": true}}), &[AnthropicBeta::Compact20260904])] #[case::context_management_edits( json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919"}]}}), - &[beta::CONTEXT_MANAGEMENT_2025_06_27] + &[AnthropicBeta::ContextManagement20250627] )] #[case::per_message_output_config( json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}), - &[beta::PER_TURN_CONTROL_2026_07_01] + &[AnthropicBeta::PerTurnControl20260701] )] #[case::advisor_tool( - json!({"tools": [{"type": ANTHROPIC_ADVISOR_TOOL_TYPE, "name": "advisor", "model": MODEL}]}), - &[beta::ADVISOR_TOOL_2026_03_01] + json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": MODEL}]}), + &[AnthropicBeta::AdvisorTool20260301] )] #[case::several_features_at_once( json!({"speed": "fast", "output_format": {"type": "json_schema"}}), - &[beta::STRUCTURED_OUTPUT, beta::FAST_MODE_2026_02_01] + &[AnthropicBeta::StructuredOutputs20251113, AnthropicBeta::FastMode20260201] )] #[tokio::test] async fn feature_betas_join_the_callers_betas_in_one_sorted_header( call: MessagesCall, #[case] fields: Value, - #[case] features: &[&str], + #[case] features: &[AnthropicBeta], ) { let upstream = upstream([message_response()]).await; let capabilities = AnthropicModelCapabilities { @@ -316,12 +299,11 @@ async fn feature_betas_join_the_callers_betas_in_one_sorted_header( .await; let sent = sent_betas(&only_request(&upstream).await); - let mut expected: Vec = features + let expected: BetaSet = features .iter() - .map(|feature| feature.to_string()) - .chain(["caller-beta-2025-01-01".to_string()]) + .cloned() + .chain([AnthropicBeta::Other("caller-beta-2025-01-01".to_string())]) .collect(); - expected.sort(); assert_eq!(sent, expected); } @@ -342,7 +324,10 @@ async fn an_oauth_key_sends_the_browser_access_header_and_the_oauth_beta(call: M request.header("anthropic-dangerous-direct-browser-access"), Some("true") ); - assert_eq!(sent_betas(&request), [ANTHROPIC_OAUTH_BETA_HEADER]); + assert_eq!( + sent_betas(&request), + BetaSet::from_iter([AnthropicBeta::Oauth20250420]) + ); assert_eq!(request.header("x-api-key"), None); } @@ -406,7 +391,6 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i custom_llm_provider: call.custom_llm_provider.clone(), extra_headers: None, provider_specific_header: None, - model: call.model.clone(), timeout: call.timeout, }, fields.clone(), @@ -664,10 +648,9 @@ async fn the_provider_prefix_is_stripped_exactly_once( let upstream = upstream([message_response()]).await; run_message(MessagesCall { - model: model.into(), api_key: Some("sk".into()), api_base: Some(upstream.uri()), - ..call + ..with_model(call, model) }) .await; diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index 133b7d2b162..14c8eb6b7c6 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -1,4 +1,7 @@ -use litellm_core::messages::{messages, types::MessagesRequest}; +use litellm_core::{ + Phase, + messages::{MessagesResponse, messages, messages_body}, +}; use litellm_http::transport::Error as TransportError; use rstest::rstest; @@ -154,7 +157,7 @@ async fn an_unreadable_success_body_is_an_invalid_response( .err() .expect("an unreadable body fails"); - assert!(error.is_response(), "{error:?}"); + assert_eq!(error.phase(), Phase::AfterSend, "{error:?}"); } #[rstest] @@ -175,47 +178,46 @@ async fn a_provider_slower_than_the_timeout_fails_the_call(call: MessagesCall) { assert!(matches!(error, Error::Transport(_)), "{error:?}"); } -fn facade_request(body: Value, api_base: &str) -> MessagesRequest<'_> { - MessagesRequest { - model: MODEL, - body, - api_key: Some("sk-ant"), - api_base: Some(api_base), - custom_llm_provider: Some("anthropic"), - extra_headers: None, - provider_specific_header: None, - timeout: Some(Duration::from_secs(5)), - shaping: MessagesShaping::default(), - } -} - +#[rstest] #[tokio::test] -async fn the_facade_runs_the_route_in_process() { +async fn the_facade_sends_through_the_injected_http_pool_configuration(call: MessagesCall) { let upstream = upstream([message_response()]).await; let base = upstream.uri(); + let settings = HttpSettings { + user_agent: Some("host-owned/1".into()), + ..HttpSettings::default() + }; - let message = messages(facade_request( - json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}), - &base, - )) + let response = messages( + &support::resources(), + &Resolution::from(&settings).config, + &RecordingSecrets::empty(), + MessagesCall { + api_key: Some("sk-ant".into()), + api_base: Some(base), + ..call + }, + ) .await .expect("messages request succeeds"); + let MessagesResponse::Message(message) = response else { + panic!("a non-streaming request returns a message"); + }; assert_eq!(message.id, "msg_1"); - assert_eq!( - only_request(&upstream).await.header("x-api-key"), - Some("sk-ant") - ); + let sent = only_request(&upstream).await; + assert_eq!(sent.header("x-api-key"), Some("sk-ant")); + assert_eq!(sent.header("user-agent"), Some("host-owned/1")); } -#[tokio::test] -async fn the_facade_rejects_a_body_that_is_not_an_object() { - let error = messages(facade_request(json!([]), UNREACHABLE_BASE)) - .await - .expect_err("a non-object body is rejected"); +#[rstest] +#[case::mistyped_param(json!({"model": MODEL, "messages": [], "max_tokens": "16"}))] +#[case::missing_messages(json!({"model": MODEL, "max_tokens": 16}))] +fn a_body_that_does_not_parse_is_an_invalid_request(#[case] raw: Value) { + let error = messages_body(object(raw)).expect_err("the body is rejected"); - assert_eq!( - error, - Error::InvalidRequest("messages body must be an object".into()) + assert!( + matches!(&error, Error::InvalidRequest(message) if message.starts_with("invalid Anthropic messages request: ")), + "{error:?}" ); } diff --git a/litellm-rust/crates/core/tests/messages/stream.rs b/litellm-rust/crates/core/tests/messages/stream.rs index c4be3127d66..f0e55eca8dd 100644 --- a/litellm-rust/crates/core/tests/messages/stream.rs +++ b/litellm-rust/crates/core/tests/messages/stream.rs @@ -1,12 +1,21 @@ -use std::{convert::Infallible, sync::Mutex}; +use std::{ + convert::Infallible, + sync::{Mutex, mpsc}, +}; use bytes::Bytes; -use litellm_core::messages::route::{Messages, MessagesStreamHead}; +use futures_util::{StreamExt, TryStreamExt}; +use litellm_core::messages::{ + MessagesResponse, messages, + route::{Messages, MessagesStreamHead}, +}; use litellm_host::host::{Demand, Host}; +use litellm_tracing::{Logger, Metadata, Record, Sink}; use rstest::rstest; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, net::TcpListener, + task::JoinHandle, }; use super::*; @@ -23,6 +32,20 @@ enum Seen { Deliver(Bytes), } +struct TraceSink(mpsc::Sender<(String, Value)>); + +impl Sink for TraceSink { + fn enabled(&self, metadata: &Metadata<'_>) -> bool { + metadata.target().starts_with("litellm_core::messages") + } + + fn emit(&self, record: &Record) { + self.0 + .send((record.message.clone(), Value::Object(record.fields.clone()))) + .unwrap(); + } +} + /// Projects like `LocalMessagesHost`, records every stream op in the order the route /// performs it, and detaches after `detach_after` ops. struct RecordingStreamHost { @@ -69,13 +92,10 @@ impl Host for RecordingStreamHost { } fn streaming(call: MessagesCall, api_base: String) -> MessagesCall { - let mut body = call.body.clone(); - body.insert("stream".into(), json!(true)); MessagesCall { api_key: Some("sk-ant".into()), api_base: Some(api_base), - body, - ..call + ..with_fields(call, json!({"stream": true})) } } @@ -87,7 +107,7 @@ fn sse_response() -> ResponseTemplate { } async fn stream_through(host: &RecordingStreamHost) -> Result { - litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await + litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await } #[rstest] @@ -123,6 +143,37 @@ async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: Me assert_eq!(delivered, SSE_BODY.as_bytes()); } +#[rstest] +#[tokio::test] +async fn debug_trace_keeps_provider_input_and_every_stream_chunk(call: MessagesCall) { + let upstream = upstream([sse_response()]).await; + let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX); + let (sender, receiver) = mpsc::channel(); + + Logger::new(TraceSink(sender)) + .instrument(stream_through(&host)) + .await + .unwrap(); + + let records: Vec<(String, Value)> = receiver.try_iter().collect(); + let request = records + .iter() + .find(|(message, _)| message == "provider request") + .unwrap(); + let body: Value = serde_json::from_str(request.1["body"].as_str().unwrap()).unwrap(); + assert_eq!(body["messages"][0]["content"], "hi"); + assert_eq!(request.1["stream"], true); + let chunks: String = records + .iter() + .filter(|(message, fields)| { + message == "stream chunk" && fields["stage"] == "provider_response" + }) + .map(|(_, fields)| fields["chunk"].as_str().unwrap()) + .collect(); + assert_eq!(chunks, SSE_BODY); + assert!(!format!("{records:?}").contains("sk-ant")); +} + #[rstest] #[case::at_open(1)] #[case::after_the_first_chunk(2)] @@ -195,10 +246,10 @@ async fn a_stream_that_ends_without_message_stop_is_relayed_as_is(call: Messages } /// Serves one SSE chunk and then holds the connection open without ever finishing. -async fn stalling_upstream() -> String { +async fn stalling_upstream() -> (String, JoinHandle<()>) { let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); let base = format!("http://{}", listener.local_addr().unwrap()); - tokio::spawn(async move { + let connection = tokio::spawn(async move { let (mut socket, _) = listener.accept().await.unwrap(); let mut request = vec![0; 4096]; let _ = socket.read(&mut request).await; @@ -209,15 +260,15 @@ async fn stalling_upstream() -> String { ) .await .unwrap(); - std::future::pending::<()>().await; + let _ = socket.read_to_end(&mut Vec::new()).await; }); - base + (base, connection) } #[rstest] #[tokio::test] async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) { - let base = stalling_upstream().await; + let (base, connection) = stalling_upstream().await; let host = RecordingStreamHost::new( MessagesCall { timeout: Some(Duration::from_millis(300)), @@ -239,11 +290,150 @@ async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) { "the chunk before the stall reached the caller, saw {} ops", seen.len() ); + tokio::time::timeout(Duration::from_secs(5), connection) + .await + .expect("timing out closes the upstream connection") + .unwrap(); +} + +#[rstest] +#[case::anthropic("anthropic")] +#[case::azure_ai("azure_ai")] +#[tokio::test] +async fn the_sdk_returns_stream_headers_and_every_sse_byte( + call: MessagesCall, + #[case] provider: &str, +) { + let upstream = upstream([sse_response()]).await; + let response = messages( + &support::resources(), + &http_config(), + &RecordingSecrets::empty(), + MessagesCall { + custom_llm_provider: Some(provider.into()), + ..streaming(call, upstream.uri()) + }, + ) + .await + .unwrap(); + + let MessagesResponse::Stream { headers, chunks } = response else { + panic!("a streaming request returns a stream"); + }; + for (name, value) in UPSTREAM_HEADERS { + assert!(headers.contains(&(name.into(), value.into()))); + } + let delivered = chunks.try_collect::>().await.unwrap().concat(); + assert_eq!(delivered, SSE_BODY.as_bytes()); + assert_eq!(only_request(&upstream).await.json()["stream"], true); } #[rstest] #[tokio::test] -async fn streaming_is_refused_for_providers_that_cannot_stream(call: MessagesCall) { +async fn the_sdk_returns_http_errors_before_opening_a_stream(call: MessagesCall) { + let upstream = upstream([ResponseTemplate::new(429).set_body_string("slow down")]).await; + let error = messages( + &support::resources(), + &http_config(), + &RecordingSecrets::empty(), + streaming(call, upstream.uri()), + ) + .await + .err() + .expect("upstream failure is returned by messages()"); + + assert_eq!( + error, + Error::Transport(litellm_http::transport::Error::Http { + status: 429, + body: "slow down".into(), + }) + ); +} + +#[rstest] +#[case::before_reading(false)] +#[case::after_reading(true)] +#[tokio::test] +async fn dropping_the_sdk_stream_closes_the_unfinished_upstream( + call: MessagesCall, + #[case] read_chunk: bool, +) { + let (base, connection) = stalling_upstream().await; + let response = tokio::time::timeout( + Duration::from_secs(5), + messages( + &support::resources(), + &http_config(), + &RecordingSecrets::empty(), + MessagesCall { + timeout: Some(Duration::from_secs(30)), + ..streaming(call, base) + }, + ), + ) + .await + .expect("messages() returns before the upstream finishes") + .unwrap(); + + let MessagesResponse::Stream { mut chunks, .. } = response else { + panic!("a streaming request returns a stream"); + }; + if read_chunk { + let chunk = tokio::time::timeout(Duration::from_secs(5), chunks.next()) + .await + .expect("the first chunk arrives before the upstream finishes") + .unwrap() + .unwrap(); + assert_eq!(chunk.as_ref(), b"event: message_start\ndata: {}\n\n"); + } + assert!(!connection.is_finished()); + drop(chunks); + tokio::time::timeout(Duration::from_secs(5), connection) + .await + .expect("dropping the stream closes the upstream connection") + .unwrap(); +} + +#[rstest] +#[tokio::test] +async fn the_sdk_yields_a_body_error_once_after_delivered_chunks(call: MessagesCall) { + let (base, connection) = stalling_upstream().await; + let response = messages( + &support::resources(), + &http_config(), + &RecordingSecrets::empty(), + MessagesCall { + timeout: Some(Duration::from_millis(300)), + ..streaming(call, base) + }, + ) + .await + .unwrap(); + + let MessagesResponse::Stream { mut chunks, .. } = response else { + panic!("a streaming request returns a stream"); + }; + assert_eq!( + chunks.next().await.unwrap().unwrap().as_ref(), + b"event: message_start\ndata: {}\n\n" + ); + let error = tokio::time::timeout(Duration::from_secs(5), chunks.next()) + .await + .expect("the stalled body times out") + .unwrap() + .unwrap_err(); + assert!(matches!(error, Error::Transport(_)), "{error:?}"); + assert!(chunks.next().await.is_none()); + tokio::time::timeout(Duration::from_secs(5), connection) + .await + .expect("the failed stream closes its upstream connection") + .unwrap(); +} + +#[rstest] +#[tokio::test] +async fn a_host_on_anthropic_sse_is_relayed_byte_for_byte(call: MessagesCall) { let upstream = upstream([sse_response()]).await; let host = RecordingStreamHost::new( MessagesCall { @@ -253,14 +443,17 @@ async fn streaming_is_refused_for_providers_that_cannot_stream(call: MessagesCal usize::MAX, ); - let error = stream_through(&host) - .await - .err() - .expect("azure streaming is refused"); + let outcome = stream_through(&host).await.expect("azure streams"); - assert_eq!( - error, - Error::Unsupported("streaming messages for this provider") - ); - assert!(received(&upstream).await.is_empty()); + assert!(matches!(outcome, MessagesOutput::Streamed)); + let seen = host.seen.into_inner().unwrap(); + let delivered: Vec = seen + .iter() + .filter_map(|step| match step { + Seen::Deliver(chunk) => Some(chunk.to_vec()), + Seen::Open(_) => None, + }) + .flatten() + .collect(); + assert_eq!(delivered, SSE_BODY.as_bytes()); } diff --git a/litellm-rust/crates/core/tests/ocr/main.rs b/litellm-rust/crates/core/tests/ocr/main.rs index 1a915389b20..e1f6b8cb5c1 100644 --- a/litellm-rust/crates/core/tests/ocr/main.rs +++ b/litellm-rust/crates/core/tests/ocr/main.rs @@ -4,6 +4,7 @@ use litellm_core::ocr::{ types::LiteLLMOcrRequest, wire::{OcrWireRequest, decode_request}, }; +use litellm_http::Client; use litellm_llms::base_llm::ocr::{ error::Error, handler::OcrClient, @@ -37,11 +38,7 @@ fn object(value: Value) -> Map { } fn ocr_client() -> OcrClient { - let document_http = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("test document client builds"); - OcrClient::for_test(reqwest::Client::new(), document_http) + OcrClient::for_test(Client::plain_for_test(), Client::no_redirect_for_test()) } async fn perform(request: LiteLLMOcrRequest) -> Result { diff --git a/litellm-rust/crates/core/tests/ocr/mistral.rs b/litellm-rust/crates/core/tests/ocr/mistral.rs index f80e564b03f..d542eeaf03a 100644 --- a/litellm-rust/crates/core/tests/ocr/mistral.rs +++ b/litellm-rust/crates/core/tests/ocr/mistral.rs @@ -1,10 +1,6 @@ use std::sync::Arc; -use litellm_auth_gcp::VertexAuth; -use litellm_http::{ - HttpClientPool, HttpSettings, Resolution, - media::{PublicDnsResolver, UrlPolicy}, -}; +use litellm_http::{HttpSettings, Resolution, media::UrlPolicy}; use litellm_llms::{ base_llm::ocr::{ settings::OcrSettings, @@ -184,6 +180,7 @@ async fn missing_credentials_come_from_the_injected_secret_source( ); } +#[rstest] #[tokio::test] async fn the_client_uses_the_injected_http_pool_configuration() { let upstream = upstream([pages_response()]).await; @@ -191,15 +188,18 @@ async fn the_client_uses_the_injected_http_pool_configuration() { user_agent: Some("host-owned/1".into()), ..HttpSettings::default() }; - let client = OcrClient::new( - &HttpClientPool::new(Arc::new(PublicDnsResolver)), - &Resolution::from(&settings).config, - UrlPolicy::default(), - VertexAuth::default(), - OcrSettings::default(), - Arc::new(litellm_secrets::source::EnvironmentSecrets::default()), - ) - .unwrap(); + let client = resources() + .ocr_client( + &Resolution::from(&settings).config, + UrlPolicy::default(), + OcrSettings::default(), + Arc::new( + litellm_secrets::source::EnvironmentSecrets::python_compatible( + litellm_http::Client::plain_for_test(), + ), + ), + ) + .unwrap(); litellm_core::ocr::client::perform( &client, diff --git a/litellm-rust/crates/core/tests/resources.rs b/litellm-rust/crates/core/tests/resources.rs new file mode 100644 index 00000000000..9764e50de1b --- /dev/null +++ b/litellm-rust/crates/core/tests/resources.rs @@ -0,0 +1,157 @@ +mod support; + +use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, +}; + +use litellm_auth::AuthServices; +use litellm_auth_gcp::{ + CredentialSource, VertexAuth, VertexAuthFuture, VertexProviderLoader, VertexTokenSource, +}; +use litellm_core::{ + ocr::{ + client::perform, + wire::{OcrWireRequest, decode_request}, + }, + resources::CoreResources, +}; +use litellm_http::{HttpSettings, Resolution}; +use litellm_llms::base_llm::ocr::settings::OcrSettings; +use rstest::{fixture, rstest}; +use serde_json::json; +use support::{ReceivedRequest, RecordingSecrets, http_pool, json_response, upstream}; + +struct TokenSource(String); + +impl VertexTokenSource for TokenSource { + fn project_id(&self) -> VertexAuthFuture<'_, String> { + Box::pin(async { Ok(self.0.clone()) }) + } + + fn token(&self) -> VertexAuthFuture<'_, String> { + Box::pin(async { Ok(self.0.clone()) }) + } +} + +#[derive(Default)] +struct Loader(AtomicUsize); + +impl VertexProviderLoader for Loader { + fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc> { + Box::pin(async move { + self.0.fetch_add(1, Ordering::SeqCst); + let identity = match source { + CredentialSource::Trusted(secret) => secret.expose().to_string(), + other => panic!("unexpected credential source: {other:?}"), + }; + Ok(Arc::new(TokenSource(identity)) as Arc) + }) + } +} + +#[fixture] +fn loader() -> Arc { + Arc::new(Loader::default()) +} + +#[fixture] +fn resources(loader: Arc) -> CoreResources { + CoreResources { + auth: Arc::new(AuthServices { + gcp: VertexAuth::new(loader), + ..AuthServices::default() + }), + pool: Arc::new(http_pool()), + } +} + +#[rstest] +#[case::shared_identity(false, "first-identity", 1)] +#[case::different_identity(false, "second-identity", 2)] +#[case::independent_resources(true, "first-identity", 2)] +#[tokio::test] +async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets( + loader: Arc, + #[with(loader.clone())] resources: CoreResources, + #[case] independent: bool, + #[case] second_identity: &str, + #[case] expected_loads: usize, +) { + let response = json_response(json!({"pages": [{"index": 0, "markdown": "hello"}]})); + let upstream = upstream([response.clone(), response]).await; + let second_resources = if independent { + CoreResources { + auth: Arc::new(AuthServices { + gcp: VertexAuth::new(loader.clone()), + ..AuthServices::default() + }), + ..resources.clone() + } + } else { + resources.clone() + }; + for (owner, identity, agent, location) in [ + (&resources, "first-identity", "first-agent", "us-central1"), + ( + &second_resources, + second_identity, + "second-agent", + "europe-west4", + ), + ] { + let http = Resolution::from(&HttpSettings { + user_agent: Some(agent.into()), + ..HttpSettings::default() + }) + .config; + let client = owner + .ocr_client( + &http, + Default::default(), + OcrSettings { + vertex_location: Some(location.into()), + ..OcrSettings::default() + }, + Arc::new(RecordingSecrets::new([("VERTEXAI_CREDENTIALS", identity)])), + ) + .unwrap(); + let request = decode_request(OcrWireRequest { + model: "vertex_ai/mistral-ocr-maas".into(), + document: json!({"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}), + api_key: None, + api_base: Some(upstream.uri()), + custom_llm_provider: None, + extra_headers: None, + optional_params: Default::default(), + input_sources: Default::default(), + timeout_seconds: Some(5.0), + }).unwrap(); + let result = perform(&client, request).await.unwrap(); + assert!(!result.pages.is_empty()); + } + let requests = upstream.received_requests().await.unwrap(); + assert_eq!(requests.len(), 2); + for (request, identity, agent, location) in [ + (&requests[0], "first-identity", "first-agent", "us-central1"), + ( + &requests[1], + second_identity, + "second-agent", + "europe-west4", + ), + ] { + assert_eq!( + request.header("authorization"), + Some(format!("Bearer {identity}").as_str()) + ); + assert_eq!(request.header("user-agent"), Some(agent)); + assert!( + request + .url + .path() + .contains(&format!("/projects/{identity}/locations/{location}/")) + ); + } + assert_eq!(loader.0.load(Ordering::SeqCst), expected_loads); +} diff --git a/litellm-rust/crates/core/tests/support/mod.rs b/litellm-rust/crates/core/tests/support/mod.rs index 4d2fe0232d0..5443437df09 100644 --- a/litellm-rust/crates/core/tests/support/mod.rs +++ b/litellm-rust/crates/core/tests/support/mod.rs @@ -3,9 +3,12 @@ #![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset -use std::sync::Mutex; +use std::sync::{Arc, Mutex}; use futures_util::future::BoxFuture; +use litellm_http::{ + HttpClientConfig, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver, +}; use litellm_secrets::{SecretValue, source::SecretSource}; use serde_json::Value; use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any}; @@ -13,6 +16,18 @@ use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any}; /// A port nothing listens on, for calls that must fail before any request is sent. pub const UNREACHABLE_BASE: &str = "http://127.0.0.1:1"; +pub fn http_pool() -> HttpClientPool { + HttpClientPool::new(Arc::new(PublicDnsResolver)) +} + +pub fn resources() -> litellm_core::resources::CoreResources { + litellm_core::resources::CoreResources::new(Arc::new(http_pool())) +} + +pub fn http_config() -> HttpClientConfig { + Resolution::from(&HttpSettings::default()).config +} + /// Starts an upstream that answers its n-th request with the n-th response and 404s after. pub async fn upstream(responses: impl IntoIterator) -> MockServer { let server = MockServer::start().await; diff --git a/litellm-rust/crates/gateway-auth/Cargo.toml b/litellm-rust/crates/gateway-auth/Cargo.toml new file mode 100644 index 00000000000..340f7224618 --- /dev/null +++ b/litellm-rust/crates/gateway-auth/Cargo.toml @@ -0,0 +1,21 @@ +[package] +name = "litellm-gateway-auth" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +axum.workspace = true +litellm-auth-types.workspace = true +litellm-config.workspace = true +litellm-secrets.workspace = true +sha2.workspace = true +subtle.workspace = true +thiserror.workspace = true + +[dev-dependencies] +futures-util.workspace = true +rstest.workspace = true +tokio.workspace = true +tower = { version = "0.5.3", features = ["util"] } diff --git a/litellm-rust/crates/gateway-auth/src/error.rs b/litellm-rust/crates/gateway-auth/src/error.rs new file mode 100644 index 00000000000..735eb7741ad --- /dev/null +++ b/litellm-rust/crates/gateway-auth/src/error.rs @@ -0,0 +1,24 @@ +use axum::{ + http::StatusCode, + response::{IntoResponse, Response}, +}; + +#[derive(Debug, thiserror::Error)] +pub enum Error { + #[error("gateway auth not configured")] + Unconfigured, + #[error("missing or invalid bearer token")] + InvalidToken, + #[error("gateway authentication unavailable")] + Secret(#[from] litellm_secrets::Error), +} + +impl IntoResponse for Error { + fn into_response(self) -> Response { + let status = match &self { + Self::InvalidToken => StatusCode::UNAUTHORIZED, + Self::Unconfigured | Self::Secret(_) => StatusCode::INTERNAL_SERVER_ERROR, + }; + (status, self.to_string()).into_response() + } +} diff --git a/litellm-rust/crates/gateway-auth/src/lib.rs b/litellm-rust/crates/gateway-auth/src/lib.rs new file mode 100644 index 00000000000..7a1fb244654 --- /dev/null +++ b/litellm-rust/crates/gateway-auth/src/lib.rs @@ -0,0 +1,74 @@ +mod error; + +use std::sync::Arc; + +use axum::{ + extract::FromRequestParts, + http::{header::AUTHORIZATION, request::Parts}, +}; +use litellm_auth_types::SecretValue; +use litellm_config::Config; +use litellm_secrets::source::SecretSource; +use sha2::{Digest, Sha256}; +use subtle::ConstantTimeEq; + +pub use error::Error; + +#[derive(Clone)] +pub struct Auth { + master_key: Option, + secrets: Arc, +} + +impl Auth { + pub fn from_config(config: &Config, secrets: Arc) -> Self { + Self { + master_key: config.general_settings.master_key.clone(), + secrets, + } + } + + async fn master_key(&self) -> Result { + let configured = self.master_key.as_ref().ok_or(Error::Unconfigured)?; + let resolved = match configured.expose().strip_prefix("os.environ/") { + Some(name) if !name.is_empty() => self + .secrets + .get_secret_str(name) + .await? + .ok_or(Error::Unconfigured)?, + Some(_) => return Err(Error::Unconfigured), + None => configured.clone(), + }; + if resolved.expose().trim().is_empty() { + return Err(Error::Unconfigured); + } + Ok(resolved) + } +} + +pub fn hash_token(token: &str) -> String { + format!("{:x}", Sha256::digest(token.as_bytes())) +} + +pub struct RequireMasterKey; + +impl FromRequestParts for RequireMasterKey { + type Rejection = Error; + + async fn from_request_parts(parts: &mut Parts, state: &Auth) -> Result { + let expected = state.master_key().await?; + let provided = parts + .headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.strip_prefix("Bearer ")) + .map(str::trim) + .ok_or(Error::InvalidToken)?; + let actual_hash = Sha256::digest(provided.as_bytes()); + let expected_hash = Sha256::digest(expected.expose().as_bytes()); + match bool::from(actual_hash.ct_eq(&expected_hash)) { + true => Ok(Self), + false => Err(Error::InvalidToken), + } + } +} diff --git a/litellm-rust/crates/gateway-auth/tests/auth.rs b/litellm-rust/crates/gateway-auth/tests/auth.rs new file mode 100644 index 00000000000..58625296664 --- /dev/null +++ b/litellm-rust/crates/gateway-auth/tests/auth.rs @@ -0,0 +1,100 @@ +use std::sync::Arc; + +use axum::{ + Router, + body::{Body, to_bytes}, + http::{Request, StatusCode}, + middleware::from_extractor_with_state, + routing::get, +}; +use futures_util::future::BoxFuture; +use litellm_auth_types::SecretValue; +use litellm_config::Config; +use litellm_gateway_auth::{Auth, RequireMasterKey, hash_token}; +use litellm_secrets::source::SecretSource; +use rstest::{fixture, rstest}; +use tower::ServiceExt; + +struct Secrets; + +impl SecretSource for Secrets { + fn get_secret_str<'a>( + &'a self, + name: &'a str, + ) -> BoxFuture<'a, Result, litellm_secrets::Error>> { + Box::pin(async move { + match name { + "MASTER_KEY" => Ok(Some(SecretValue::new("resolved-key"))), + "EMPTY" => Ok(Some(SecretValue::new(""))), + "ERROR" => Err(litellm_secrets::Error::ExternalRead(Box::new( + std::io::Error::other("private-backend-detail"), + ))), + _ => Ok(None), + } + }) + } +} + +#[fixture] +fn secrets() -> Arc { + Arc::new(Secrets) +} + +#[rstest] +#[case::literal("literal-key", Some("Bearer literal-key"), 204)] +#[case::reference("os.environ/MASTER_KEY", Some("Bearer resolved-key"), 204)] +#[case::reference_is_not_a_token( + "os.environ/MASTER_KEY", + Some("Bearer os.environ/MASTER_KEY"), + 401 +)] +#[case::wrong("literal-key", Some("Bearer other-key"), 401)] +#[case::missing("literal-key", None, 401)] +#[case::wrong_scheme("literal-key", Some("Basic literal-key"), 401)] +#[case::empty_token("literal-key", Some("Bearer "), 401)] +#[case::missing_reference("os.environ/MISSING", Some("Bearer os.environ/MISSING"), 500)] +#[case::empty_reference("os.environ/EMPTY", Some("Bearer "), 500)] +#[case::empty_key("", Some("Bearer "), 500)] +#[case::whitespace_key(" ", Some("Bearer "), 500)] +#[case::empty_reference_name("os.environ/", Some("Bearer os.environ/"), 500)] +#[case::secret_failure("os.environ/ERROR", Some("Bearer private-backend-detail"), 500)] +#[tokio::test] +async fn enforces_configured_keys_without_exposing_secrets( + secrets: Arc, + #[case] key: &str, + #[case] authorization: Option<&str>, + #[case] status: u16, +) { + let config = Config::from_yaml(&format!( + "model_list: []\ngeneral_settings:\n master_key: '{key}'\n" + )) + .unwrap(); + let app = Router::new() + .route("/protected", get(|| async { StatusCode::NO_CONTENT })) + .layer(from_extractor_with_state::( + Auth::from_config(&config, secrets), + )); + let request = Request::get("/protected"); + let request = match authorization { + Some(value) => request.header("authorization", value), + None => request, + }; + let response = app + .oneshot(request.body(Body::empty()).unwrap()) + .await + .unwrap(); + assert_eq!(response.status().as_u16(), status); + let body = to_bytes(response.into_body(), 4096).await.unwrap(); + let text = std::str::from_utf8(&body).unwrap(); + assert!(!text.contains("private-backend-detail")); + assert!(!text.contains("literal-key")); + assert!(!text.contains("resolved-key")); +} + +#[rstest] +fn hash_token_matches_python_sha256_hexdigest() { + assert_eq!( + hash_token("sk-1234"), + "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b" + ); +} diff --git a/litellm-rust/crates/gateway-inference/AGENTS.md b/litellm-rust/crates/gateway-inference/AGENTS.md new file mode 100644 index 00000000000..7dc57380083 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/AGENTS.md @@ -0,0 +1,5 @@ +- Expose a mountable Axum router; listener binding, server lifecycle, and shared inbound middleware belong to `gateway` +- Own the public inference HTTP boundary: endpoint paths, request parsing, model alias resolution, response envelopes, and SSE delivery +- Delegate inference execution to `core` and provider transformations and authentication to `llms` and the auth crates; do not duplicate them in handlers +- Use injected deployments, HTTP pools, settings, and secret sources; do not load process configuration or construct independent clients in handlers +- Test HTTP contracts here, including status codes, forwarded headers, error envelopes, and streaming behavior; keep core and provider tests in their owning crates diff --git a/litellm-rust/crates/gateway-inference/Cargo.toml b/litellm-rust/crates/gateway-inference/Cargo.toml new file mode 100644 index 00000000000..f5ee5e81ba6 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/Cargo.toml @@ -0,0 +1,28 @@ +[package] +name = "litellm-gateway-inference" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +axum = { workspace = true, features = ["json", "multipart"] } +base64.workspace = true +bytes.workspace = true +futures-util.workspace = true +litellm-auth.workspace = true +litellm-core.workspace = true +litellm-http.workspace = true +litellm-llms.workspace = true +litellm-router.workspace = true +litellm-secrets.workspace = true +litellm-types.workspace = true +serde_json.workspace = true +thiserror.workspace = true + +[dev-dependencies] +futures-util.workspace = true +tokio = { workspace = true, features = ["io-util"] } +rstest.workspace = true +tower = { version = "0.5.3", features = ["util"] } +wiremock = "0.6.5" diff --git a/litellm-rust/crates/gateway-inference/src/audio_transcription.rs b/litellm-rust/crates/gateway-inference/src/audio_transcription.rs new file mode 100644 index 00000000000..d5fd6603e20 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/audio_transcription.rs @@ -0,0 +1,59 @@ +use std::{path::Path, sync::Arc}; + +use axum::{ + Json, + extract::{Request, State}, + response::{IntoResponse, Response}, +}; +use base64::{Engine, engine::general_purpose::STANDARD}; +use litellm_core::audio_transcription::{audio_transcription, types::AudioTranscriptionRequest}; +use serde_json::{Value, json}; + +use crate::{Error, Gateway, request}; + +pub(crate) async fn create(State(gateway): State>, request: Request) -> Response { + match handle(&gateway, request).await { + Ok(response) => Json(response).into_response(), + Err(error) => error.openai_response(), + } +} + +async fn handle(gateway: &Gateway, request: Request) -> Result { + let (body, upload) = request::parse(request).await?; + let deployment = request::deployment(gateway, &body)?; + let audio = match upload { + Some(upload) => { + let format = upload + .file_name + .as_deref() + .and_then(|name| Path::new(name).extension()) + .and_then(|extension| extension.to_str()) + .ok_or_else(|| { + Error::InvalidBody("audio file requires a filename extension".into()) + })?; + json!({"data": STANDARD.encode(upload.bytes), "format": format.to_ascii_lowercase()}) + } + None => body + .get("audio") + .cloned() + .ok_or_else(|| Error::InvalidBody("audio is required".into()))?, + }; + Ok(audio_transcription( + &gateway.resources, + &gateway.http, + AudioTranscriptionRequest { + model: &deployment.model, + audio, + api_key: deployment.api_key.as_deref(), + api_base: deployment.api_base.as_deref(), + custom_llm_provider: deployment.custom_llm_provider.as_deref(), + extra_headers: None, + optional_params: body + .into_iter() + .filter(|(name, _)| !matches!(name.as_str(), "model" | "audio")) + .collect(), + timeout: deployment.timeout, + }, + ) + .await?) +} diff --git a/litellm-rust/crates/gateway-inference/src/chat_completions.rs b/litellm-rust/crates/gateway-inference/src/chat_completions.rs new file mode 100644 index 00000000000..c386f38fe06 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/chat_completions.rs @@ -0,0 +1,83 @@ +use std::sync::Arc; + +use axum::{ + Json, + body::Bytes, + extract::{Path, State}, + http::StatusCode, + response::{IntoResponse, Response}, +}; +use litellm_core::chat_completions::{chat_completions, types::ChatCompletionsRequest}; +use serde_json::{Map, Value}; + +use crate::{Error, Gateway, request}; + +pub(crate) async fn create(State(gateway): State>, body: Bytes) -> Response { + respond(&gateway, request::object(&body)).await +} + +pub(crate) async fn deployment( + State(gateway): State>, + Path(path): Path, + body: Bytes, +) -> Response { + if let Some(model) = path + .strip_suffix("/chat/completions") + .filter(|model| !model.is_empty()) + { + let body = request::object(&body).map(|body| { + if body.get("model").is_some_and(|model| !model.is_null()) { + return body; + } + body.into_iter() + .chain([("model".into(), Value::String(model.into()))]) + .collect() + }); + return respond(&gateway, body).await; + } + if path.ends_with("/embeddings") || path.ends_with("/completions") { + return Error::Unsupported(path).openai_response(); + } + StatusCode::NOT_FOUND.into_response() +} + +async fn respond(gateway: &Gateway, body: Result, Error>) -> Response { + let result = match body { + Ok(body) => handle(gateway, body).await, + Err(error) => Err(error), + }; + match result { + Ok(response) => response, + Err(error) => error.openai_response(), + } +} + +async fn handle(gateway: &Gateway, body: Map) -> Result { + let deployment = request::deployment(gateway, &body)?; + if body.get("stream").and_then(Value::as_bool) == Some(true) { + return Err(Error::Unsupported("streaming chat completions".into())); + } + let messages = body + .get("messages") + .cloned() + .ok_or_else(|| Error::InvalidBody("messages is required".into()))?; + let response = chat_completions( + &gateway.resources, + &gateway.http, + ChatCompletionsRequest { + model: &deployment.model, + messages, + optional_params: body + .into_iter() + .filter(|(name, _)| !matches!(name.as_str(), "model" | "messages" | "stream")) + .collect(), + api_key: deployment.api_key.as_deref(), + api_base: deployment.api_base.as_deref(), + custom_llm_provider: deployment.custom_llm_provider.as_deref(), + extra_headers: None, + timeout: deployment.timeout, + }, + ) + .await?; + Ok(Json(response).into_response()) +} diff --git a/litellm-rust/crates/gateway-inference/src/error.rs b/litellm-rust/crates/gateway-inference/src/error.rs new file mode 100644 index 00000000000..1d40857d962 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/error.rs @@ -0,0 +1,225 @@ +use axum::http::StatusCode; +use axum::{ + Json, + response::{IntoResponse, Response}, +}; +use litellm_core::RouteError; +use litellm_http::transport::Error as TransportError; +use litellm_llms::base_llm::ocr::error::Error as OcrError; +use serde_json::{Map, Value, json}; + +#[derive(Debug, thiserror::Error)] +pub enum Error { + #[error("invalid request body: {0}")] + InvalidBody(String), + #[error( + "/v1/messages: Invalid model name passed in model={0}. Call `/v1/models` to view available models for your key." + )] + UnknownModel(String), + #[error(transparent)] + Route(#[from] RouteError), + #[error(transparent)] + Ocr(#[from] OcrError), + #[error("{0} is not implemented by the Rust gateway")] + Unsupported(String), + #[error("request body exceeds the size limit")] + BodyTooLarge, + #[error("{0}")] + Internal(String), +} + +impl Error { + pub fn status(&self) -> StatusCode { + match self { + Self::Unsupported(_) + | Self::Route(RouteError::Unsupported(_)) + | Self::Ocr(OcrError::Unsupported(_)) => StatusCode::NOT_IMPLEMENTED, + Self::BodyTooLarge => StatusCode::PAYLOAD_TOO_LARGE, + Self::Ocr( + OcrError::Auth(litellm_auth::Error::MissingApiKey { .. }) + | OcrError::MissingAzureAiCredentials + | OcrError::MissingAzureDocumentIntelligenceCredentials + | OcrError::MissingReductoApiKey, + ) => StatusCode::UNAUTHORIZED, + Self::Ocr(error) => error + .http_status_code() + .and_then(|status| StatusCode::from_u16(status).ok()) + .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), + Self::InvalidBody(_) | Self::UnknownModel(_) => StatusCode::BAD_REQUEST, + Self::Route(RouteError::Transport(TransportError::Http { status, .. })) => { + StatusCode::from_u16(*status).unwrap_or(StatusCode::BAD_GATEWAY) + } + Self::Route(RouteError::Auth(litellm_auth::Error::MissingApiKey { .. })) => { + StatusCode::UNAUTHORIZED + } + Self::Route(error) if error.is_request() => StatusCode::BAD_REQUEST, + Self::Route(_) | Self::Internal(_) => StatusCode::INTERNAL_SERVER_ERROR, + } + } + + pub fn openai_response(self) -> Response { + let status = self.status(); + let message = match &self { + Self::UnknownModel(model) => format!("Invalid model name passed in model={model}"), + _ => self.to_string(), + }; + ( + status, + Json(json!({"error": { + "message": message, + "type": error_type(status), + "param": null, + "code": status.as_u16(), + }})), + ) + .into_response() + } + + /// The Anthropic error envelope Python's `AnthropicExceptionMapping` builds: an upstream + /// body already in that shape passes through, any other has its message extracted. + pub fn body(&self, request_id: Option<&str>) -> Value { + let raw = match self { + Self::Route(RouteError::Transport(TransportError::Http { body, .. })) => body.clone(), + other => other.to_string(), + }; + let parsed = serde_json::from_str::(&raw).ok(); + let envelope = match parsed { + Some(Value::Object(object)) if is_anthropic_error(&object) => object, + Some(Value::Object(object)) => { + envelope(self.status(), provider_message(&object).unwrap_or(&raw)) + } + _ => envelope(self.status(), &raw), + }; + Value::Object(with_request_id(envelope, request_id)) + } + + /// An `event: error` frame, for a stream that fails after its headers went out. + pub fn sse_frame(&self) -> String { + format!("event: error\ndata: {}\n\n", self.body(None)) + } +} + +fn error_type(status: StatusCode) -> &'static str { + match status.as_u16() { + 400 => "invalid_request_error", + 401 => "authentication_error", + 403 => "permission_error", + 404 => "not_found_error", + 413 => "request_too_large", + 429 => "rate_limit_error", + 529 => "overloaded_error", + _ => "api_error", + } +} + +fn envelope(status: StatusCode, message: &str) -> Map { + let Value::Object(envelope) = json!({ + "type": "error", + "error": {"type": error_type(status), "message": message}, + }) else { + unreachable!("a json object literal is an object") + }; + envelope +} + +fn is_anthropic_error(object: &Map) -> bool { + object.get("type").and_then(Value::as_str) == Some("error") + && object + .get("error") + .and_then(Value::as_object) + .is_some_and(|error| error.contains_key("type") && error.contains_key("message")) +} + +fn provider_message(object: &Map) -> Option<&str> { + if let Some(detail) = object.get("detail").and_then(Value::as_object) { + return detail.get("message").and_then(Value::as_str); + } + ["Message", "message"] + .into_iter() + .filter_map(|key| object.get(key).and_then(Value::as_str)) + .find(|message| !message.is_empty()) +} + +fn with_request_id(envelope: Map, request_id: Option<&str>) -> Map { + match request_id { + Some(id) if !id.is_empty() && !envelope.contains_key("request_id") => envelope + .into_iter() + .chain([("request_id".to_string(), Value::from(id))]) + .collect(), + _ => envelope, + } +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + + use super::*; + + fn upstream(status: u16, body: &str) -> Error { + Error::Route(RouteError::Transport(TransportError::Http { + status, + body: body.into(), + })) + } + + #[rstest] + #[case::anthropic_body_passes_through( + upstream(529, r#"{"type":"error","error":{"type":"overloaded_error","message":"busy","extra":1}}"#), + Some("req_1"), + json!({"type": "error", "error": {"type": "overloaded_error", "message": "busy", "extra": 1}, "request_id": "req_1"}), + )] + #[case::upstream_request_id_wins( + upstream(400, r#"{"type":"error","error":{"type":"x","message":"m"},"request_id":"upstream"}"#), + Some("caller"), + json!({"type": "error", "error": {"type": "x", "message": "m"}, "request_id": "upstream"}), + )] + #[case::bedrock_detail( + upstream(403, r#"{"detail":{"message":"denied"}}"#), + None, + json!({"type": "error", "error": {"type": "permission_error", "message": "denied"}}), + )] + #[case::aws_message( + upstream(429, r#"{"Message":"slow down"}"#), + None, + json!({"type": "error", "error": {"type": "rate_limit_error", "message": "slow down"}}), + )] + #[case::plain_text_with_unmapped_status( + upstream(502, "bad gateway"), + None, + json!({"type": "error", "error": {"type": "api_error", "message": "bad gateway"}}), + )] + #[case::unknown_model( + Error::UnknownModel("nope".into()), + None, + json!({"type": "error", "error": { + "type": "invalid_request_error", + "message": "/v1/messages: Invalid model name passed in model=nope. Call `/v1/models` to view available models for your key.", + }}), + )] + fn body_follows_the_anthropic_exception_mapping( + #[case] error: Error, + #[case] request_id: Option<&str>, + #[case] expected: Value, + ) { + assert_eq!(error.body(request_id), expected); + } + + #[rstest] + #[case::upstream_status(upstream(429, ""), StatusCode::TOO_MANY_REQUESTS)] + #[case::rejected_request(Error::Route(RouteError::InvalidRequest("top_k".into())), StatusCode::BAD_REQUEST)] + #[case::missing_key( + Error::Route(RouteError::Auth(litellm_auth::Error::MissingApiKey { + provider: "Anthropic", + environment_variable: "ANTHROPIC_API_KEY", + })), + StatusCode::UNAUTHORIZED, + )] + #[case::lost_connection( + Error::Route(RouteError::Transport(TransportError::Network("reset".into()))), + StatusCode::INTERNAL_SERVER_ERROR, + )] + fn status_follows_who_is_at_fault(#[case] error: Error, #[case] status: StatusCode) { + assert_eq!(error.status(), status); + } +} diff --git a/litellm-rust/crates/gateway-inference/src/lib.rs b/litellm-rust/crates/gateway-inference/src/lib.rs new file mode 100644 index 00000000000..eebe3f34a09 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/lib.rs @@ -0,0 +1,59 @@ +//! The proxy's inference endpoints as an axum [`Router`] a server mounts. +//! +//! Authentication, rate limiting and logging are the mounting server's layers; this crate +//! maps a public model name to its deployment and runs the core route. + +mod audio_transcription; +mod chat_completions; +mod error; +pub mod messages; +mod ocr; +mod request; + +use std::sync::Arc; + +use axum::{Router, routing::post}; +use litellm_core::resources::CoreResources; +use litellm_http::HttpClientConfig; +use litellm_llms::base_llm::ocr::handler::OcrClient; +use litellm_secrets::source::SecretSource; + +pub use error::Error; +pub use litellm_router::{Deployment, Router as ModelList}; + +pub struct Gateway { + pub resources: CoreResources, + pub http: HttpClientConfig, + pub secrets: Arc, + pub models: ModelList, + pub ocr: OcrClient, +} + +pub fn router(gateway: Arc) -> Router { + Router::new() + .route("/v1/messages", post(messages::create)) + .route("/ocr", post(ocr::create)) + .route("/v1/ocr", post(ocr::create)) + .route("/chat/completions", post(chat_completions::create)) + .route("/v1/chat/completions", post(chat_completions::create)) + .route("/engines/{*path}", post(chat_completions::deployment)) + .route( + "/openai/deployments/{*path}", + post(chat_completions::deployment), + ) + .route("/audio/transcriptions", post(audio_transcription::create)) + .route( + "/v1/audio/transcriptions", + post(audio_transcription::create), + ) + .route("/responses", post(request::unsupported)) + .route("/v1/responses", post(request::unsupported)) + .route("/embeddings", post(request::unsupported)) + .route("/v1/embeddings", post(request::unsupported)) + .route("/completions", post(request::unsupported)) + .route("/v1/completions", post(request::unsupported)) + .layer(axum::extract::DefaultBodyLimit::max( + request::MAX_BODY_BYTES, + )) + .with_state(gateway) +} diff --git a/litellm-rust/crates/gateway-inference/src/messages/mod.rs b/litellm-rust/crates/gateway-inference/src/messages/mod.rs new file mode 100644 index 00000000000..5d6a8faa0e8 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/messages/mod.rs @@ -0,0 +1,122 @@ +//! `POST /v1/messages`, as the Python proxy's `anthropic_response` serves it. + +use std::{convert::Infallible, sync::Arc}; + +use axum::{ + Json, + body::{Body, Bytes}, + extract::State, + http::{HeaderMap, StatusCode, header}, + response::{IntoResponse, Response}, +}; +use futures_util::{StreamExt, stream::BoxStream}; +use litellm_core::messages::{ + Error as RouteError, MessagesCall, MessagesResponse, messages, messages_body, +}; +use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders}; +use serde_json::{Map, Value}; + +use crate::{Deployment, Error, Gateway}; + +/// Client headers Python forwards to Anthropic-speaking providers on every call. +const ANTHROPIC_API_HEADERS: [&str; 2] = ["anthropic-version", "anthropic-beta"]; +const ANTHROPIC_API_HEADER_PROVIDERS: &str = "anthropic,bedrock,bedrock_mantle,vertex_ai"; + +pub async fn create( + State(gateway): State>, + headers: HeaderMap, + body: Bytes, +) -> Response { + let request_id = headers + .get("x-request-id") + .and_then(|value| value.to_str().ok()) + .map(str::to_owned); + match handle(&gateway, &headers, &body).await { + Ok(response) => response, + Err(error) => (error.status(), Json(error.body(request_id.as_deref()))).into_response(), + } +} + +async fn handle(gateway: &Gateway, headers: &HeaderMap, body: &[u8]) -> Result { + let body = match serde_json::from_slice(body) { + Ok(Value::Object(body)) => body, + Ok(_) => return Err(Error::InvalidBody("expected a JSON object".into())), + Err(error) => return Err(Error::InvalidBody(error.to_string())), + }; + let model_name = body + .get("model") + .and_then(Value::as_str) + .ok_or_else(|| Error::InvalidBody("model is required".into()))?; + let deployment = gateway + .models + .get(model_name) + .ok_or_else(|| Error::UnknownModel(model_name.to_owned()))?; + let call = project(deployment, body, headers)?; + match messages( + &gateway.resources, + &gateway.http, + gateway.secrets.as_ref(), + call, + ) + .await? + { + MessagesResponse::Message(message) => Ok(Json(message).into_response()), + MessagesResponse::Stream { chunks, .. } => Ok(stream(chunks)), + } +} + +fn project( + deployment: &Deployment, + body: Map, + headers: &HeaderMap, +) -> Result { + let body = body + .into_iter() + .map(|(name, value)| match name.as_str() { + "model" => (name, Value::from(deployment.model.as_str())), + _ => (name, value), + }) + .collect(); + Ok(MessagesCall { + body: messages_body(body)?, + api_key: deployment.api_key.clone(), + api_base: deployment.api_base.clone(), + custom_llm_provider: deployment.custom_llm_provider.clone(), + extra_headers: None, + provider_specific_header: anthropic_api_headers(headers), + timeout: deployment.timeout, + shaping: deployment.shaping.clone(), + }) +} + +fn anthropic_api_headers(headers: &HeaderMap) -> Option { + let extra_headers: Map = ANTHROPIC_API_HEADERS + .into_iter() + .filter_map(|name| { + let value = headers.get(name)?.to_str().ok()?; + Some((name.to_owned(), Value::from(value))) + }) + .collect(); + (!extra_headers.is_empty()).then(|| { + ProviderSpecificHeaders::One(ProviderSpecificHeader { + custom_llm_provider: ANTHROPIC_API_HEADER_PROVIDERS.into(), + extra_headers, + }) + }) +} + +/// A chunk that fails after the stream opened is delivered as an SSE error frame, since +/// the status line already went out; the stream ends on it. +fn stream(chunks: BoxStream<'static, Result>) -> Response { + let body = chunks.map(|chunk| { + Ok::<_, Infallible>( + chunk.unwrap_or_else(|error| Bytes::from(Error::Route(error).sse_frame())), + ) + }); + ( + StatusCode::OK, + [(header::CONTENT_TYPE, "text/event-stream")], + Body::from_stream(body), + ) + .into_response() +} diff --git a/litellm-rust/crates/gateway-inference/src/ocr.rs b/litellm-rust/crates/gateway-inference/src/ocr.rs new file mode 100644 index 00000000000..d666223e037 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/ocr.rs @@ -0,0 +1,77 @@ +use std::sync::Arc; + +use axum::{ + Json, + extract::{Request, State}, + response::{IntoResponse, Response}, +}; +use litellm_auth::SecretValue; +use litellm_core::ocr::{ + client::perform, + types::{LiteLLMOcrRequest, OcrConnectionInputs, OcrDocumentInput}, +}; +use litellm_llms::base_llm::ocr::transformation::OcrDocument; +use serde_json::Value; + +use crate::{Error, Gateway, request}; + +pub(crate) async fn create(State(gateway): State>, request: Request) -> Response { + match handle(&gateway, request).await { + Ok(response) => Json(response).into_response(), + Err(error) => error.openai_response(), + } +} + +async fn handle(gateway: &Gateway, request: Request) -> Result { + let header_format = request + .headers() + .get("x-req-format") + .and_then(|value| value.to_str().ok()) + .map(str::to_owned); + let (body, upload) = request::parse(request).await?; + let deployment = request::deployment(gateway, &body)?; + let document = match upload { + Some(upload) => OcrDocumentInput::Bytes { + bytes: upload.bytes, + file_name: upload.file_name, + mime_type: upload.mime_type, + }, + None => OcrDocument::try_from( + body.get("document") + .cloned() + .ok_or_else(|| Error::InvalidBody("document is required".into()))?, + )? + .into(), + }; + let format = body + .get("req_format") + .filter(|value| !value.is_null()) + .cloned() + .or_else(|| header_format.map(Value::String)); + let format = format.map(|value| match value { + Value::String(value) => Value::String(value.trim().to_ascii_lowercase()), + value => value, + }); + let options = body + .into_iter() + .filter(|(name, _)| !matches!(name.as_str(), "model" | "document" | "req_format")) + .chain(format.map(|value| ("req_format".into(), value))) + .collect(); + let call = LiteLLMOcrRequest::from_inputs( + deployment.model.clone(), + document, + deployment.custom_llm_provider.as_deref(), + options, + OcrConnectionInputs { + api_key: deployment.api_key.clone().map(SecretValue::new), + api_base: deployment.api_base.clone(), + timeout: deployment.timeout, + ..Default::default() + }, + )?; + let response = perform(&gateway.ocr, call).await?; + match response.provider_native_response { + Some(native) => Ok(Value::Object(native)), + None => Ok(response.into_json()), + } +} diff --git a/litellm-rust/crates/gateway-inference/src/request.rs b/litellm-rust/crates/gateway-inference/src/request.rs new file mode 100644 index 00000000000..f58c7b3ed79 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/request.rs @@ -0,0 +1,108 @@ +use axum::{ + body::{Bytes, to_bytes}, + extract::{FromRequest, Multipart, Request}, + http::Uri, + response::Response, +}; +use serde_json::{Map, Value}; + +use crate::{Deployment, Error, Gateway}; + +pub(crate) const MAX_FILE_BYTES: usize = 50 * 1024 * 1024; +pub(crate) const MAX_BODY_BYTES: usize = MAX_FILE_BYTES + 1024 * 1024; + +pub(crate) struct Upload { + pub bytes: Bytes, + pub file_name: Option, + pub mime_type: Option, +} + +pub(crate) fn object(body: &[u8]) -> Result, Error> { + match serde_json::from_slice(body) { + Ok(Value::Object(body)) => Ok(body), + Ok(_) => Err(Error::InvalidBody("expected a JSON object".into())), + Err(error) => Err(Error::InvalidBody(error.to_string())), + } +} + +pub(crate) fn deployment<'a>( + gateway: &'a Gateway, + body: &Map, +) -> Result<&'a Deployment, Error> { + let model = body + .get("model") + .and_then(Value::as_str) + .ok_or_else(|| Error::InvalidBody("model is required".into()))?; + gateway + .models + .get(model) + .ok_or_else(|| Error::UnknownModel(model.to_owned())) +} + +pub(crate) async fn parse(request: Request) -> Result<(Map, Option), Error> { + let multipart = request + .headers() + .get("content-type") + .and_then(|header| header.to_str().ok()) + .is_some_and(|value| { + value + .to_ascii_lowercase() + .starts_with("multipart/form-data") + }); + if !multipart { + let body = to_bytes(request.into_body(), MAX_BODY_BYTES) + .await + .map_err(|_| Error::BodyTooLarge)?; + return Ok((object(&body)?, None)); + } + let mut multipart = Multipart::from_request(request, &()) + .await + .map_err(|error| Error::InvalidBody(error.to_string()))?; + let mut fields = Map::new(); + let mut upload = None; + while let Some(field) = multipart.next_field().await.map_err(multipart_error)? { + let name = field.name().unwrap_or_default().to_owned(); + if name == "file" { + let file_name = field.file_name().map(str::to_owned); + let mime_type = field + .content_type() + .and_then(|value| value.split(';').next()) + .map(str::trim) + .filter(|value| *value != "application/octet-stream") + .map(str::to_owned); + let bytes = field.bytes().await.map_err(multipart_error)?; + if bytes.len() > MAX_FILE_BYTES { + return Err(Error::BodyTooLarge); + } + if bytes.is_empty() { + return Err(Error::InvalidBody("uploaded file is empty".into())); + } + upload = Some(Upload { + bytes, + file_name, + mime_type, + }); + } else if name != "document" { + let text = field.text().await.map_err(multipart_error)?; + let value = serde_json::from_str(&text).unwrap_or(Value::String(text)); + fields.insert(name, value); + } + } + if upload.is_none() { + return Err(Error::InvalidBody( + "multipart request requires a file field".into(), + )); + } + Ok((fields, upload)) +} + +fn multipart_error(error: axum::extract::multipart::MultipartError) -> Error { + if error.status() == axum::http::StatusCode::PAYLOAD_TOO_LARGE { + return Error::BodyTooLarge; + } + Error::InvalidBody(error.to_string()) +} + +pub(crate) async fn unsupported(uri: Uri) -> Response { + Error::Unsupported(uri.path().to_owned()).openai_response() +} diff --git a/litellm-rust/crates/gateway-inference/tests/messages.rs b/litellm-rust/crates/gateway-inference/tests/messages.rs new file mode 100644 index 00000000000..30836498d49 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/tests/messages.rs @@ -0,0 +1,121 @@ +mod support; + +use axum::{ + body::{Body, to_bytes}, + http::Request, +}; +use rstest::rstest; +use serde_json::json; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tower::ServiceExt; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_json, header, method, path}, +}; + +#[rstest] +#[case(false)] +#[case(true)] +#[tokio::test] +async fn messages_reaches_the_provider_and_preserves_json_or_sse(#[case] streaming: bool) { + let upstream = MockServer::start().await; + let message = json!({"id": "msg_test", "type": "message", "role": "assistant", + "model": "test-model", "content": [{"type": "text", "text": "hello"}], + "stop_reason": "end_turn", "usage": {"input_tokens": 1, "output_tokens": 1}}); + let sse = "event: message_start\ndata: {\"type\":\"message_start\"}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + let template = if streaming { + ResponseTemplate::new(200).set_body_raw(sse, "text/event-stream") + } else { + ResponseTemplate::new(200).set_body_json(message.clone()) + }; + let messages = json!([{"role": "user", "content": "hi"}]); + Mock::given(method("POST")).and(path("/v1/messages")) + .and(header("x-api-key", "test-key")) + .and(header("anthropic-beta", "test-feature")) + .and(body_json(json!({"model": "test-model", "messages": messages, "max_tokens": 16, "stream": streaming}))) + .respond_with(template).expect(1).mount(&upstream).await; + let request = Request::post("/v1/messages") + .header("content-type", "application/json").header("anthropic-beta", "test-feature") + .body(Body::from(json!({"model": "public/model", "messages": messages, "max_tokens": 16, "stream": streaming}).to_string())).unwrap(); + let response = support::app("anthropic/test-model", &upstream.uri()) + .oneshot(request) + .await + .unwrap(); + assert_eq!(response.status(), 200); + if streaming { + assert_eq!(response.headers()["content-type"], "text/event-stream"); + assert_eq!(to_bytes(response.into_body(), 4096).await.unwrap(), sse); + } else { + let body = support::json(response).await; + assert_eq!(body["content"], message["content"]); + assert_eq!(body["usage"], message["usage"]); + } +} + +#[tokio::test] +async fn invalid_messages_stays_an_anthropic_error() { + let response = support::post( + support::app("anthropic/test-model", "http://127.0.0.1:1"), + "/v1/messages", + json!({"model": "public/model", "messages": "invalid", "max_tokens": 16}), + ) + .await; + assert_eq!(response.status(), 400); + let body = support::json(response).await; + assert_eq!(body["type"], "error"); + assert_eq!(body["error"]["type"], "invalid_request_error"); +} + +/// Answers with the SSE head and one event, then drops the connection short of the +/// announced body length. +async fn truncating_upstream() -> String { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let base = format!("http://{}", listener.local_addr().unwrap()); + tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut request = vec![0; 4096]; + let _ = socket.read(&mut request).await; + socket + .write_all( + format!( + "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ncontent-length: {}\r\n\r\n{FIRST_EVENT}", + FIRST_EVENT.len() * 2 + ) + .as_bytes(), + ) + .await + .unwrap(); + }); + base +} + +const FIRST_EVENT: &str = "event: message_start\ndata: {}\n\n"; + +#[tokio::test] +async fn a_stream_that_fails_after_opening_ends_with_an_sse_error_frame() { + let base = truncating_upstream().await; + let request = Request::post("/v1/messages") + .header("content-type", "application/json") + .body(Body::from( + json!({"model": "public/model", "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 16, "stream": true}) + .to_string(), + )) + .unwrap(); + + let response = support::app("anthropic/test-model", &base) + .oneshot(request) + .await + .unwrap(); + + assert_eq!(response.status(), 200); + let body = to_bytes(response.into_body(), 4096).await.unwrap(); + let text = std::str::from_utf8(&body).unwrap(); + let frame = text + .strip_prefix(FIRST_EVENT) + .and_then(|rest| rest.strip_prefix("event: error\ndata: ")) + .unwrap_or_else(|| panic!("the delivered event then one error frame, got {text:?}")); + let error: serde_json::Value = serde_json::from_str(frame.trim_end()).unwrap(); + assert_eq!(error["type"], "error"); + assert_eq!(error["error"]["type"], "api_error"); +} diff --git a/litellm-rust/crates/gateway-inference/tests/ocr.rs b/litellm-rust/crates/gateway-inference/tests/ocr.rs new file mode 100644 index 00000000000..3d5a2d22b09 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/tests/ocr.rs @@ -0,0 +1,118 @@ +mod support; + +use axum::{body::Body, http::Request}; +use rstest::rstest; +use serde_json::{Value, json}; +use tower::ServiceExt; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_json, header, method, path}, +}; + +const DOCUMENT: &str = "data:application/pdf;base64,YWJj"; + +#[rstest] +#[case("/ocr", false)] +#[case("/v1/ocr", true)] +#[tokio::test] +async fn json_and_multipart_reach_ocr_with_the_deployment( + #[case] route: &str, + #[case] multipart: bool, +) { + let upstream = MockServer::start().await; + Mock::given(method("POST")).and(path("/v1/ocr")) + .and(header("authorization", "Bearer test-key")) + .and(body_json(json!({"model": "test-ocr", "document": {"type": "document_url", "document_url": DOCUMENT}, "pages": [0]}))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"pages": [{"index": 0, "markdown": "recognized text"}]}))) + .expect(1).mount(&upstream).await; + let app = support::app("mistral/test-ocr", &upstream.uri()); + let response = if multipart { + let body = "--boundary\r\nContent-Disposition: form-data; name=\"model\"\r\n\r\npublic/model\r\n--boundary\r\nContent-Disposition: form-data; name=\"pages\"\r\n\r\n[0]\r\n--boundary\r\nContent-Disposition: form-data; name=\"file\"; filename=\"test.pdf\"\r\nContent-Type: application/pdf\r\n\r\nabc\r\n--boundary--\r\n"; + app.oneshot( + Request::post(route) + .header("content-type", "multipart/form-data; boundary=boundary") + .body(Body::from(body)) + .unwrap(), + ) + .await + .unwrap() + } else { + support::post(app, route, json!({"model": "public/model", "document": {"type": "document_url", "document_url": DOCUMENT}, "pages": [0]})).await + }; + assert_eq!(response.status(), 200); + let body = support::json(response).await; + assert_eq!(body["pages"][0]["markdown"], "recognized text"); + assert_eq!(body["model"], "test-ocr"); +} + +#[rstest] +#[case(None, true)] +#[case(Some("litellm"), false)] +#[tokio::test] +async fn native_format_header_is_used_unless_the_body_overrides_it( + #[case] format: Option<&str>, + #[case] native: bool, +) { + let upstream = MockServer::start().await; + let payload = json!({"pages": [{"index": 0, "markdown": "text"}], "provider_only": true}); + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json(payload.clone())) + .expect(1) + .mount(&upstream) + .await; + let body = json!({"model": "public/model", "document": {"type": "document_url", "document_url": DOCUMENT}, "req_format": format}); + let response = support::app("mistral/test-ocr", &upstream.uri()) + .oneshot( + Request::post("/ocr") + .header("x-req-format", " Native ") + .body(Body::from(body.to_string())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), 200); + let body = support::json(response).await; + if native { + assert_eq!(body, payload); + } else { + assert_eq!(body["object"], "ocr"); + assert_eq!( + body["pages"][0]["markdown"], + payload["pages"][0]["markdown"] + ); + } +} + +#[tokio::test] +async fn ocr_keeps_upstream_status_in_an_openai_error_envelope() { + let upstream = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(429).set_body_json(json!({"message": "busy"}))) + .expect(1) + .mount(&upstream) + .await; + let response = support::post(support::app("mistral/test-ocr", &upstream.uri()), "/ocr", + json!({"model": "public/model", "document": {"type": "document_url", "document_url": DOCUMENT}})).await; + assert_eq!(response.status(), 429); + let body = support::json(response).await; + assert_eq!(body["error"]["code"], 429); + assert!(body["error"]["message"].as_str().unwrap().contains("busy")); +} + +#[rstest] +#[case(json!({"model": "public/model"}))] +#[case(json!({"model": "public/model", "document": "/etc/passwd"}))] +#[case(json!({"model": "missing", "document": {"type": "document_url", "document_url": DOCUMENT}}))] +#[tokio::test] +async fn invalid_ocr_requests_do_not_call_the_provider(#[case] body: Value) { + let upstream = MockServer::start().await; + let response = support::post( + support::app("mistral/test-ocr", &upstream.uri()), + "/ocr", + body, + ) + .await; + assert_eq!(response.status(), 400); + assert!(support::json(response).await["error"]["message"].is_string()); + assert!(upstream.received_requests().await.unwrap().is_empty()); +} diff --git a/litellm-rust/crates/gateway-inference/tests/routes.rs b/litellm-rust/crates/gateway-inference/tests/routes.rs new file mode 100644 index 00000000000..5d3b06ccadb --- /dev/null +++ b/litellm-rust/crates/gateway-inference/tests/routes.rs @@ -0,0 +1,92 @@ +mod support; + +use rstest::rstest; +use serde_json::{Value, json}; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_partial_json, method}, +}; + +#[rstest] +#[case("/chat/completions", Some("public/model"))] +#[case("/v1/chat/completions", Some("public/model"))] +#[case("/engines/public/model/chat/completions", None)] +#[case("/openai/deployments/public/model/chat/completions", None)] +#[case("/openai/deployments/unused/chat/completions", Some("public/model"))] +#[tokio::test] +async fn chat_aliases_call_core_and_use_the_body_model_before_the_path( + #[case] route: &str, + #[case] model: Option<&str>, +) { + let upstream = MockServer::start().await; + let messages = json!([{"role": "user", "content": "hi"}]); + Mock::given(method("POST")) + .and(body_partial_json( + json!({"model": "test-model", "max_tokens": 16}), + )) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "id": "msg_test", "model": "test-model", "content": [{"type": "text", "text": "hello"}], + "stop_reason": "end_turn", "usage": {"input_tokens": 1, "output_tokens": 1} + }))) + .expect(1) + .mount(&upstream) + .await; + let response = support::post( + support::app("anthropic/test-model", &upstream.uri()), + route, + json!({"model": model, "messages": messages, "max_tokens": 16}), + ) + .await; + assert_eq!(response.status(), 200); + assert_eq!( + support::json(response).await["choices"][0]["message"]["content"], + "hello" + ); +} + +#[rstest] +#[case("/responses")] +#[case("/v1/responses")] +#[case("/embeddings")] +#[case("/v1/embeddings")] +#[case("/completions")] +#[case("/v1/completions")] +#[case("/engines/public/model/embeddings")] +#[case("/openai/deployments/public/model/completions")] +#[tokio::test] +async fn unimplemented_routes_return_an_explicit_error(#[case] path: &str) { + let response = support::post( + support::app("anthropic/test-model", "http://127.0.0.1:1"), + path, + json!({}), + ) + .await; + assert_eq!(response.status(), 501); + assert!( + support::json(response).await["error"]["message"] + .as_str() + .unwrap() + .contains("not implemented") + ); +} + +#[rstest] +#[case("/audio/transcriptions")] +#[case("/v1/audio/transcriptions")] +#[tokio::test] +async fn transcription_aliases_reach_core_validation(#[case] path: &str) { + let response = support::post( + support::app("bedrock/test-model", "http://127.0.0.1:1"), + path, + json!({"model": "public/model", "audio": {"data": "YWJj", "format": "invalid"}}), + ) + .await; + assert_eq!(response.status(), 400); + let body: Value = support::json(response).await; + assert!( + body["error"]["message"] + .as_str() + .unwrap() + .contains("audio.format") + ); +} diff --git a/litellm-rust/crates/gateway-inference/tests/support/mod.rs b/litellm-rust/crates/gateway-inference/tests/support/mod.rs new file mode 100644 index 00000000000..d56489d28cd --- /dev/null +++ b/litellm-rust/crates/gateway-inference/tests/support/mod.rs @@ -0,0 +1,75 @@ +use std::{sync::Arc, time::Duration}; + +use axum::{ + Router, + body::{Body, to_bytes}, + http::Request, + response::Response, +}; +use futures_util::future::BoxFuture; +use litellm_core::resources::CoreResources; +use litellm_gateway_inference::{Deployment, Gateway, router}; +use litellm_http::{HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver}; +use litellm_llms::base_llm::ocr::settings::OcrSettings; +use litellm_secrets::{SecretValue, source::SecretSource}; +use serde_json::Value; +use tower::ServiceExt; + +struct NoSecrets; + +impl SecretSource for NoSecrets { + fn get_secret_str<'a>( + &'a self, + _: &'a str, + ) -> BoxFuture<'a, Result, litellm_secrets::Error>> { + Box::pin(async { Ok(None) }) + } +} + +pub fn app(model: &str, api_base: &str) -> Router { + let pool = Arc::new(HttpClientPool::new(Arc::new(PublicDnsResolver))); + let http = Resolution::from(&HttpSettings::default()).config; + let secrets = Arc::new(NoSecrets); + let resources = CoreResources::new(pool); + let ocr = resources + .ocr_client( + &http, + Default::default(), + OcrSettings::default(), + secrets.clone(), + ) + .unwrap(); + router(Arc::new(Gateway { + resources, + http, + secrets, + ocr, + models: [( + "public/model".into(), + Deployment { + model: model.into(), + api_base: Some(api_base.into()), + api_key: Some("test-key".into()), + timeout: Some(Duration::from_secs(5)), + ..Default::default() + }, + )] + .into_iter() + .collect(), + })) +} + +pub async fn post(app: Router, path: &str, body: Value) -> Response { + app.oneshot( + Request::post(path) + .header("content-type", "application/json") + .body(Body::from(body.to_string())) + .unwrap(), + ) + .await + .unwrap() +} + +pub async fn json(response: Response) -> Value { + serde_json::from_slice(&to_bytes(response.into_body(), 1024 * 1024).await.unwrap()).unwrap() +} diff --git a/litellm-rust/crates/gateway/AGENTS.md b/litellm-rust/crates/gateway/AGENTS.md new file mode 100644 index 00000000000..10a42b9ec3b --- /dev/null +++ b/litellm-rust/crates/gateway/AGENTS.md @@ -0,0 +1,5 @@ +- Keep this crate a thin composition layer: mount endpoint routers and serve the supplied listener +- Server lifecycle and shared inbound middleware belong here, including client authentication, rate limiting, and request logging +- Endpoint paths, request handling, model resolution, and response encoding belong to the mounted crates; provider execution belongs to `core` and `llms` +- Inject shared state and infrastructure; avoid global runtimes, duplicate client pools, and abstractions for hypothetical endpoint groups +- Test mounting and server lifecycle through public HTTP behavior; test endpoint semantics in the owning crate diff --git a/litellm-rust/crates/gateway/Cargo.toml b/litellm-rust/crates/gateway/Cargo.toml new file mode 100644 index 00000000000..c27a3f5b17e --- /dev/null +++ b/litellm-rust/crates/gateway/Cargo.toml @@ -0,0 +1,28 @@ +[package] +name = "litellm-gateway" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +axum.workspace = true +http-body-util = "0.1" +litellm-core.workspace = true +litellm-gateway-inference.workspace = true +litellm-gateway-auth.workspace = true +litellm-config.workspace = true +litellm-http.workspace = true +litellm-llms.workspace = true +litellm-secrets.workspace = true +litellm-tracing.workspace = true +serde_json.workspace = true +tracing.workspace = true +tokio.workspace = true +uuid.workspace = true + +[dev-dependencies] +futures-util.workspace = true +rstest.workspace = true +tokio = { workspace = true, features = ["sync"] } +tower = { version = "0.5", features = ["util"] } diff --git a/litellm-rust/crates/gateway/src/lib.rs b/litellm-rust/crates/gateway/src/lib.rs new file mode 100644 index 00000000000..fc16dc0de67 --- /dev/null +++ b/litellm-rust/crates/gateway/src/lib.rs @@ -0,0 +1,174 @@ +use std::{sync::Arc, time::Instant}; + +use axum::{ + Router, + body::{Body, Bytes}, + extract::Request, + middleware::Next, + response::Response, +}; +use http_body_util::BodyExt; + +use litellm_config::Config; +use litellm_core::resources::CoreResources; +use litellm_gateway_auth::{Auth, RequireMasterKey}; +use litellm_gateway_inference::{Gateway, ModelList}; +use litellm_http::{ + ClientVariant, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver, +}; +use litellm_llms::base_llm::ocr::settings::OcrSettings; +use litellm_secrets::source::EnvironmentSecrets; +use litellm_tracing::ByteChunk; +use uuid::Uuid; + +pub fn build_inference(config: &Config) -> Result, litellm_http::Error> { + let pool = Arc::new(HttpClientPool::new(Arc::new(PublicDnsResolver))); + let http = Resolution::from(&HttpSettings::default()).config; + let client = pool.client(&http, ClientVariant::Provider)?; + let secrets = Arc::new(EnvironmentSecrets::python_compatible(client)); + let resources = CoreResources::new(pool); + let ocr = resources.ocr_client( + &http, + Default::default(), + OcrSettings::default(), + secrets.clone(), + )?; + + Ok(Arc::new(Gateway { + resources, + http, + secrets, + models: ModelList::from_model_list(&config.model_list), + ocr, + })) +} + +pub fn router(inference: Arc, config: &Config) -> Router { + let auth = Auth::from_config(config, inference.secrets.clone()); + litellm_gateway_inference::router(inference) + .route_layer(axum::middleware::from_extractor_with_state::< + RequireMasterKey, + _, + >(auth)) + .layer(axum::middleware::from_fn(log_request)) +} + +async fn log_request(request: Request, next: Next) -> Response { + let request_id = Uuid::new_v4().to_string(); + let log_body_chunks = tracing::enabled!(tracing::Level::DEBUG); + let method = request.method().clone(); + let path = request.uri().path().to_owned(); + let started = Instant::now(); + let request = if log_body_chunks { + request.map(|body| logged_body(body, request_id.clone(), "input")) + } else { + request + }; + let response = next.run(request).await; + tracing::info!( + %request_id, + %method, + %path, + status = response.status().as_u16(), + time_to_headers_ms = started.elapsed().as_secs_f64() * 1000.0, + "response headers" + ); + if log_body_chunks { + response.map(|body| logged_body(body, request_id, "output")) + } else { + response + } +} + +fn logged_body(body: Body, request_id: String, direction: &'static str) -> Body { + Body::new(body.map_frame(move |frame| { + if let Some(data) = frame.data_ref() { + log_chunk(&request_id, direction, data); + } + frame + })) +} + +fn log_chunk(request_id: &str, direction: &str, data: &Bytes) { + let chunk = ByteChunk::new(data); + tracing::debug!(request_id, direction, encoding = chunk.encoding(), chunk = %chunk, "body chunk"); +} + +#[cfg(test)] +mod tests { + use std::{convert::Infallible, sync::mpsc}; + + use axum::{body::to_bytes, http::StatusCode, routing::post}; + use futures_util::stream; + use litellm_tracing::{Logger, Metadata, Record, Sink}; + use rstest::rstest; + use serde_json::{Value, json}; + use tower::ServiceExt; + + use super::*; + + struct LogSink(mpsc::Sender); + + impl Sink for LogSink { + fn enabled(&self, _: &Metadata<'_>) -> bool { + true + } + + fn emit(&self, record: &Record) { + self.0 + .send(json!({"message": record.message, "fields": record.fields})) + .unwrap(); + } + } + + #[rstest] + #[tokio::test] + async fn logs_each_body_chunk_without_changing_streamed_bytes() { + let app = Router::new() + .route( + "/stream", + post(|_: Bytes| async { + ( + StatusCode::OK, + Body::from_stream(stream::iter([ + Ok::<_, Infallible>(Bytes::from_static(b"event: first\n\n")), + Ok(Bytes::from_static(b"event: second\n\n")), + ])), + ) + }), + ) + .layer(axum::middleware::from_fn(log_request)); + let request_chunks = [ + Ok::<_, Infallible>(Bytes::from_static(b"hello")), + Ok(Bytes::from_static(b" world")), + ]; + let request = Request::post("/stream") + .body(Body::from_stream(stream::iter(request_chunks))) + .unwrap(); + let (sender, receiver) = mpsc::channel(); + let logger = Logger::new(LogSink(sender)); + + let output = logger + .instrument(async { + let response = app.oneshot(request).await.unwrap(); + to_bytes(response.into_body(), 1024).await.unwrap() + }) + .await; + + assert_eq!(output, "event: first\n\nevent: second\n\n"); + let records: Vec = receiver.try_iter().collect(); + assert_eq!(records.len(), 5); + assert_eq!(records[0]["fields"]["chunk"], "hello"); + assert_eq!(records[1]["fields"]["chunk"], " world"); + assert_eq!(records[2]["fields"]["status"], 200); + assert_eq!(records[3]["fields"]["chunk"], "event: first\n\n"); + assert_eq!(records[4]["fields"]["chunk"], "event: second\n\n"); + let request_id = &records[2]["fields"]["request_id"]; + assert!(request_id.as_str().is_some()); + assert!( + records + .iter() + .all(|record| &record["fields"]["request_id"] == request_id) + ); + } +} diff --git a/litellm-rust/crates/gateway/src/main.rs b/litellm-rust/crates/gateway/src/main.rs new file mode 100644 index 00000000000..40f9cb442d5 --- /dev/null +++ b/litellm-rust/crates/gateway/src/main.rs @@ -0,0 +1,57 @@ +use std::{ + error::Error, + time::{SystemTime, UNIX_EPOCH}, +}; + +use litellm_config::Config; +use litellm_tracing::{Level, Logger, Metadata, Record, Sink}; +use serde_json::json; + +struct StderrSink { + level: Level, +} + +impl Sink for StderrSink { + fn enabled(&self, metadata: &Metadata<'_>) -> bool { + *metadata.level() <= self.level && metadata.target().starts_with("litellm") + } + + fn emit(&self, record: &Record) { + let timestamp_ms = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_millis(); + eprintln!( + "{}", + json!({ + "timestamp_ms": timestamp_ms, + "level": record.metadata.level().as_str(), + "target": record.metadata.target(), + "message": record.message, + "fields": record.fields, + }) + ); + } +} + +#[tokio::main] +async fn main() -> Result<(), Box> { + let level = std::env::var("RUST_LOG") + .ok() + .and_then(|value| value.parse().ok()) + .unwrap_or(Level::INFO); + Logger::new(StderrSink { level }).install_global()?; + let config_path = std::env::var("LITELLM_CONFIG").unwrap_or_else(|_| "config.yaml".into()); + let config = Config::load(config_path)?; + let inference = litellm_gateway::build_inference(&config)?; + let host = std::env::var("HOST").unwrap_or_else(|_| "0.0.0.0".into()); + let port = std::env::var("PORT") + .unwrap_or_else(|_| "4000".into()) + .parse::()?; + let listener = tokio::net::TcpListener::bind((host.as_str(), port)).await?; + + tracing::info!(address = %listener.local_addr()?, models = config.model_list.len(), log_level = %level, "gateway listening"); + + axum::serve(listener, litellm_gateway::router(inference, &config)).await?; + Ok(()) +} diff --git a/litellm-rust/crates/gateway/tests/server.rs b/litellm-rust/crates/gateway/tests/server.rs new file mode 100644 index 00000000000..a91051441ea --- /dev/null +++ b/litellm-rust/crates/gateway/tests/server.rs @@ -0,0 +1,143 @@ +use std::{ + sync::{Arc, mpsc}, + time::Duration, +}; + +use axum::{body::Body, http::Request}; +use litellm_config::Config; +use litellm_gateway_inference::{Error, Gateway}; +use litellm_http::ClientVariant; +use litellm_tracing::{Logger, Metadata, Record, Sink}; +use rstest::{fixture, rstest}; +use serde_json::{Value, json}; +use tokio::{net::TcpListener, sync::oneshot, time::timeout}; +use tower::ServiceExt; + +struct LogSink(mpsc::Sender); + +impl Sink for LogSink { + fn enabled(&self, _: &Metadata<'_>) -> bool { + true + } + + fn emit(&self, record: &Record) { + self.0 + .send(json!({"message": record.message, "fields": record.fields})) + .unwrap(); + } +} + +#[fixture] +fn inference() -> Arc { + litellm_gateway::build_inference(&Config::from_yaml("model_list: []").unwrap()).unwrap() +} + +#[rstest] +#[case::authorized("/v1/messages", Some("Bearer gateway-key"), Some("gateway-key"), 400)] +#[case::missing_token("/v1/messages", None, Some("gateway-key"), 401)] +#[case::wrong_token("/v1/messages", Some("Bearer wrong"), Some("gateway-key"), 401)] +#[case::ocr("/ocr", None, Some("gateway-key"), 401)] +#[case::chat("/v1/chat/completions", None, Some("gateway-key"), 401)] +#[case::deployment( + "/openai/deployments/model/chat/completions", + None, + Some("gateway-key"), + 401 +)] +#[case::transcription("/audio/transcriptions", None, Some("gateway-key"), 401)] +#[case::unsupported_route("/responses", None, Some("gateway-key"), 401)] +#[case::unknown_path("/unknown", None, Some("gateway-key"), 404)] +#[case::unknown_path_unconfigured("/unknown", None, None, 404)] +#[case::unconfigured("/v1/messages", Some("Bearer gateway-key"), None, 500)] +#[tokio::test] +async fn authenticates_before_serving_mounted_inference_routes( + inference: Arc, + #[case] path: &str, + #[case] authorization: Option<&str>, + #[case] master_key: Option<&str>, + #[case] status: u16, +) { + let config = Config::from_yaml(&format!( + "model_list: []\ngeneral_settings:\n master_key: {}\n", + master_key.unwrap_or("null") + )) + .unwrap(); + let client = inference + .resources + .pool + .client(&inference.http, ClientVariant::Provider) + .unwrap(); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let (shutdown, stopped) = oneshot::channel(); + let server = tokio::spawn(async move { + axum::serve(listener, litellm_gateway::router(inference, &config)) + .with_graceful_shutdown(async move { + let _ = stopped.await; + }) + .await + }); + + let request = client + .post(format!("http://{address}{path}")) + .timeout(Duration::from_secs(5)) + .header("x-request-id", "gateway-test") + .json(&json!({"model": "unconfigured-model"})); + let request = match authorization { + Some(value) => request.header("authorization", value), + None => request, + }; + let response = request.send().await.unwrap(); + assert_eq!(response.status().as_u16(), status); + if status == 400 { + let expected = Error::UnknownModel("unconfigured-model".into()); + assert_eq!( + response.json::().await.unwrap(), + expected.body(Some("gateway-test")) + ); + } else { + let text = response.text().await.unwrap(); + assert!(!text.contains("gateway-key")); + assert!(!text.contains("unconfigured-model")); + } + + shutdown.send(()).unwrap(); + timeout(Duration::from_secs(5), server) + .await + .unwrap() + .unwrap() + .unwrap(); +} + +#[rstest] +#[tokio::test] +async fn logs_request_outcome_without_credentials_or_query(inference: Arc) { + let config = + Config::from_yaml("model_list: []\ngeneral_settings:\n master_key: gateway-key\n") + .unwrap(); + let request = Request::builder() + .method("POST") + .uri("/v1/messages?token=query-secret") + .header("authorization", "Bearer header-secret") + .body(Body::empty()) + .unwrap(); + let (sender, receiver) = mpsc::channel(); + let logger = Logger::new(LogSink(sender)); + + let response = logger + .instrument(litellm_gateway::router(inference, &config).oneshot(request)) + .await + .unwrap(); + + assert_eq!(response.status().as_u16(), 401); + let record = receiver.try_recv().unwrap(); + assert_eq!(record["message"], "response headers"); + assert_eq!(record["fields"]["method"], "POST"); + assert_eq!(record["fields"]["path"], "/v1/messages"); + assert_eq!(record["fields"]["status"], 401); + assert!(record["fields"]["time_to_headers_ms"].as_f64().unwrap() >= 0.0); + assert!(record["fields"]["request_id"].as_str().is_some()); + assert!(receiver.try_recv().is_err()); + assert!(!record.to_string().contains("header-secret")); + assert!(!record.to_string().contains("query-secret")); +} diff --git a/litellm-rust/crates/host-python/AGENTS.md b/litellm-rust/crates/host-python/AGENTS.md index 7c1919f9f39..cadc55a35a7 100644 --- a/litellm-rust/crates/host-python/AGENTS.md +++ b/litellm-rust/crates/host-python/AGENTS.md @@ -2,7 +2,7 @@ - Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`ProtocolHost` traits - No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features - The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business - - `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, not the caller's dict; a protocol host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance) + - `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, rewritten in place by the route's `Preflight`, not the caller's dict; a protocol host that projects from it inherits the adapter's rewrites (for the legacy adapter: setup, deployment hooks) and the preflight's (credential inheritance) - A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is - A failing `classify` is raised with the native error's text as its `__context__`, never swallowed - Use standard PyO3 ownership and conversion APIs diff --git a/litellm-rust/crates/host-python/Cargo.toml b/litellm-rust/crates/host-python/Cargo.toml index c1b35c0f69d..bb77a1bf330 100644 --- a/litellm-rust/crates/host-python/Cargo.toml +++ b/litellm-rust/crates/host-python/Cargo.toml @@ -6,15 +6,17 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-host.workspace = true + bytes.workspace = true futures-util.workspace = true -litellm-host.workspace = true -pyo3.workspace = true -pyo3-async-runtimes.workspace = true -pythonize.workspace = true serde.workspace = true tokio = { workspace = true, features = ["rt", "sync"] } +pyo3.workspace = true +pyo3-async-runtimes.workspace = true +pythonize = "0.29.0" + [dev-dependencies] rstest.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index 7f07475bc4c..83ed6416d6e 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -9,6 +9,12 @@ pub fn missing_state() -> PyErr { PyRuntimeError::new_err("missing native call state") } +/// The SDK's request policy, run by the driver on the keyword view `begin` returned and +/// before the protocol host projects from it. It rewrites that view in place, so the +/// lifecycle that returned it sees the rewrite too; a rejection fails the call as a host +/// failure, so the lifecycle still observes it. +pub type Preflight = fn(Python<'_>, &Bound<'_, PyDict>) -> PyResult<()>; + /// What an adapter step produced: either the value the driver asked for, or a Python /// awaitable the driver hands back to the caller's task before asking again. pub enum LifecycleStep { diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 372af2843bd..50eae1e0225 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -14,7 +14,8 @@ use pyo3::types::PyDict; use tokio::sync::Mutex; use crate::adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, Preflight, ProtocolHost, PythonLifecycle, + missing_state, }; use crate::execution::{poll_async_value, run_async_value, run_sync_value}; use crate::handle::{Execution, ExecutionBody, ExecutionStep}; @@ -83,6 +84,7 @@ where { host: H, adapter: Box, + preflight: Preflight, machine: Option>>>, arguments: Option>, started_at: f64, @@ -95,12 +97,14 @@ where } /// Runs one native call for Python: synchronously, or as a coroutine that awaits every -/// host suspension inline in the caller's task. +/// host suspension inline in the caller's task. `preflight` runs once, on the keyword view +/// the adapter's `begin` returned, before the host projects from it. pub fn run_call( py: Python<'_>, machine: M, host: H, adapter: Box, + preflight: Preflight, arguments: Py, asynchronous: bool, ) -> PyResult> @@ -111,6 +115,7 @@ where let mut driver = PythonDriver { host, adapter, + preflight, machine: Some(Arc::new(Mutex::new(MachineState { machine, result: None, @@ -213,6 +218,9 @@ where match (expect, step) { (Expect::Started, LifecycleStep::Done) => self.begin(py), (Expect::Arguments, LifecycleStep::Arguments(arguments)) => { + if let Err(error) = (self.preflight)(py, arguments.bind(py)) { + return self.adapter_failed(py, error); + } self.arguments = Some(arguments); self.stage = Stage::Call; self.resume_machine(py, None) @@ -869,6 +877,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri host: SyntheticHost, script: AdapterScript, asynchronous: bool, + ) -> (PyResult>, Vec) { + run_preflighted(py, machine, host, script, no_preflight, asynchronous) + } + + fn no_preflight(_: Python<'_>, _: &Bound<'_, PyDict>) -> PyResult<()> { + Ok(()) + } + + fn run_preflighted( + py: Python<'_>, + machine: CallMachine, + host: SyntheticHost, + script: AdapterScript, + preflight: Preflight, + asynchronous: bool, ) -> (PyResult>, Vec) { let log = Log(host.log.0.clone()); let adapter = SyntheticAdapter { @@ -882,6 +905,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri machine, host, Box::new(adapter), + preflight, arguments.unbind(), asynchronous, ); @@ -1088,6 +1112,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri streaming_machine(), StreamingHost, Box::new(adapter), + no_preflight, PyDict::new(py).unbind(), asynchronous, ) @@ -1291,6 +1316,87 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } + /// The rejection a preflight raised, kept so a test can check the caller receives that + /// exact object. A `Preflight` is a plain `fn`, so it cannot capture one itself. + static REJECTION: Mutex>> = Mutex::new(None); + + fn rejecting_preflight(py: Python<'_>, _: &Bound<'_, PyDict>) -> PyResult<()> { + let error = PyValueError::new_err("over budget"); + *REJECTION.lock().unwrap() = Some(error.value(py).clone().unbind()); + Err(error) + } + + fn inheriting_preflight(_: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<()> { + arguments.set_item("api_key", "inherited") + } + + #[test] + fn a_preflight_rejection_is_the_callers_error_and_the_machine_never_starts() { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + for asynchronous in [false, true] { + let (result, log) = run_preflighted( + py, + success_machine(), + SyntheticHost { + log: Log::default(), + op: OpScript::Answer, + classifier_fails: false, + }, + AdapterScript::Plain, + rejecting_preflight, + asynchronous, + ); + let error = result.unwrap_err(); + let raised = REJECTION.lock().unwrap().take().unwrap(); + assert!(error.value(py).is(&raised)); + assert_eq!( + log, + [ + "started", + "begin", + "failed:Host:over budget", + "adapter.close", + "host.close" + ] + ); + } + }); + } + + #[test] + fn the_host_projects_from_the_keyword_view_the_preflight_rewrote() { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + for asynchronous in [false, true] { + let (result, _) = run_preflighted( + py, + success_machine(), + SyntheticHost { + log: Log::default(), + op: OpScript::Answer, + classifier_fails: false, + }, + AdapterScript::Plain, + inheriting_preflight, + asynchronous, + ); + assert_eq!( + result.unwrap().extract::(py).unwrap(), + "project:2|sign|rewritten" + ); + } + }); + } + #[test] fn the_adapters_finalized_response_is_what_the_call_returns_and_reports() { let _guard = PYTHON_GLOBALS @@ -1412,6 +1518,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri success_machine(), host, Box::new(adapter), + no_preflight, PyDict::new(py).unbind(), false, ) diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 7e17c4da51e..2f9e37fe968 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -15,7 +15,8 @@ mod handle; mod marshal; pub use adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, Preflight, ProtocolHost, PythonLifecycle, + missing_state, }; pub use argument::lookup; pub use callable::wrap_failure; diff --git a/litellm-rust/crates/host/src/hooks.rs b/litellm-rust/crates/host/src/hooks.rs new file mode 100644 index 00000000000..14b0f1ea08a --- /dev/null +++ b/litellm-rust/crates/host/src/hooks.rs @@ -0,0 +1,145 @@ +use std::future::Future; + +use crate::{ + event::{MachineEvent, RequestContext, WireRequest}, + machine::{HostChannel, MachineFault}, + protocol::Protocol, +}; + +/// What a route reaches for mid-call: the send-time rewrite and the events it reports. +/// Python's `logging_obj.pre_call` and `post_call`, in that order. +pub trait RouteHooks: Send + Sync { + fn before_send( + &self, + wire: WireRequest, + context: RequestContext, + ) -> impl Future> + Send; + + fn emit(&self, event: MachineEvent) -> impl Future> + Send; +} + +/// No host: the wire request goes out as prepared and nothing observes the call. +impl RouteHooks for () { + async fn before_send(&self, wire: WireRequest, _: RequestContext) -> Result { + Ok(wire) + } + + async fn emit(&self, _: MachineEvent) -> Result<(), E> { + Ok(()) + } +} + +impl RouteHooks for HostChannel +where + R::Error: From, +{ + async fn before_send( + &self, + wire: WireRequest, + context: RequestContext, + ) -> Result { + HostChannel::before_send(self, wire, context).await + } + + async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> { + HostChannel::emit(self, event).await + } +} + +#[cfg(test)] +mod tests { + use std::convert::Infallible; + + use serde_json::json; + + use super::*; + use crate::{ + event::RawResponse, + host::HostOp, + machine::{CallMachine, Machine, MachineStep}, + }; + + struct Unit; + + #[derive(Clone, Debug)] + struct Fault; + + impl Protocol for Unit { + type Response = (WireRequest, ()); + type Error = Fault; + type Projection = (); + type Op = Infallible; + type Chunk = Infallible; + type StreamHead = Infallible; + } + + impl From for Fault { + fn from(_: MachineFault) -> Self { + Fault + } + } + + fn wire(url: &str) -> WireRequest { + WireRequest { + url: url.into(), + headers: Vec::new(), + body: json!({}), + } + } + + fn context() -> RequestContext { + RequestContext { + model: "m".into(), + custom_llm_provider: "p".into(), + optional_params: json!({}), + secret_fields: Vec::new(), + api_key: None, + } + } + + #[tokio::test] + async fn the_channel_yields_each_hook_as_its_op_and_returns_the_answer() { + let mut machine = CallMachine::::new(|channel| { + Box::pin(async move { + let sent = RouteHooks::before_send(&channel, wire("prepared"), context()).await?; + RouteHooks::emit( + &channel, + MachineEvent::ResponseReceived { + raw: RawResponse { body: "raw".into() }, + }, + ) + .await?; + Ok((sent, ())) + }) + }); + + let Ok(MachineStep::Host(HostOp::BeforeSend { wire, reply, .. })) = machine.resume().await + else { + panic!("before_send yields BeforeSend"); + }; + assert_eq!(wire.url, "prepared"); + reply.send(WireRequest { + url: "rewritten".into(), + ..*wire + }); + + let Ok(MachineStep::Host(HostOp::Emit(event, reply))) = machine.resume().await else { + panic!("emit yields Emit"); + }; + assert!(matches!(event, MachineEvent::ResponseReceived { .. })); + reply.send(()); + + let Ok(MachineStep::Complete((sent, ()))) = machine.resume().await else { + panic!("the call completes with the answers"); + }; + assert_eq!(sent.url, "rewritten"); + } + + #[tokio::test] + async fn no_hooks_pass_the_wire_request_through() { + let sent = RouteHooks::::before_send(&(), wire("prepared"), context()) + .await + .unwrap(); + assert_eq!(sent.url, "prepared"); + } +} diff --git a/litellm-rust/crates/host/src/lib.rs b/litellm-rust/crates/host/src/lib.rs index c6b9e59b65a..1df68941fa3 100644 --- a/litellm-rust/crates/host/src/lib.rs +++ b/litellm-rust/crates/host/src/lib.rs @@ -7,6 +7,7 @@ //! may rewrite the wire request before it is sent. pub mod event; +pub mod hooks; pub mod host; pub mod machine; pub mod protocol; diff --git a/litellm-rust/crates/http/Cargo.toml b/litellm-rust/crates/http/Cargo.toml index cad5aa87e49..0cb2b15b768 100644 --- a/litellm-rust/crates/http/Cargo.toml +++ b/litellm-rust/crates/http/Cargo.toml @@ -22,5 +22,7 @@ veil.workspace = true webpki-roots.workspace = true [dev-dependencies] +rcgen = "0.14.10" +tempfile.workspace = true rstest.workspace = true tokio.workspace = true diff --git a/litellm-rust/crates/http/src/client.rs b/litellm-rust/crates/http/src/client.rs new file mode 100644 index 00000000000..1f7017d083b --- /dev/null +++ b/litellm-rust/crates/http/src/client.rs @@ -0,0 +1,38 @@ +use std::ops::Deref; + +#[derive(Clone, Debug)] +pub struct Client(reqwest::Client); + +impl Client { + pub(crate) fn new(client: reqwest::Client) -> Self { + Self(client) + } + + #[cfg(any(test, feature = "test-support"))] + pub fn for_test(client: reqwest::Client) -> Self { + Self(client) + } + + #[cfg(any(test, feature = "test-support"))] + pub fn plain_for_test() -> Self { + Self(reqwest::Client::new()) + } + + #[cfg(any(test, feature = "test-support"))] + pub fn no_redirect_for_test() -> Self { + Self( + reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("a client without TLS or proxy settings builds"), + ) + } +} + +impl Deref for Client { + type Target = reqwest::Client; + + fn deref(&self) -> &reqwest::Client { + &self.0 + } +} diff --git a/litellm-rust/crates/http/src/config.rs b/litellm-rust/crates/http/src/config.rs index cb0173369d5..2f36784bc70 100644 --- a/litellm-rust/crates/http/src/config.rs +++ b/litellm-rust/crates/http/src/config.rs @@ -18,10 +18,16 @@ pub enum Verify { BuiltInRoots, } +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +pub enum ClientIdentity { + Pem(PathBuf), + Split { certificate: PathBuf, key: PathBuf }, +} + #[derive(Clone, Debug, PartialEq, Eq, Hash)] pub struct HttpClientConfig { pub verify: Verify, - pub client_certificate: Option, + pub client_certificate: Option, pub key_exchange_group: Option, pub tls12_cipher_suites: Option>, pub force_ipv4: bool, @@ -67,7 +73,7 @@ impl From<&HttpSettings> for Resolution { Self { config: HttpClientConfig { verify: Verify::from(settings), - client_certificate: settings.ssl_certificate.clone(), + client_certificate: settings.ssl_certificate.clone().map(ClientIdentity::Pem), key_exchange_group: curve.clone().ok().flatten(), tls12_cipher_suites: ciphers.tls12_cipher_suites, force_ipv4: settings.force_ipv4, @@ -276,7 +282,7 @@ mod tests { config, HttpClientConfig { verify: Verify::BuiltInRoots, - client_certificate: Some("/client.pem".into()), + client_certificate: Some(ClientIdentity::Pem("/client.pem".into())), key_exchange_group: None, tls12_cipher_suites: None, force_ipv4: true, diff --git a/litellm-rust/crates/http/src/lib.rs b/litellm-rust/crates/http/src/lib.rs index a1456208bb3..3e55a1843c8 100644 --- a/litellm-rust/crates/http/src/lib.rs +++ b/litellm-rust/crates/http/src/lib.rs @@ -1,3 +1,10 @@ +#![allow( + clippy::disallowed_types, + clippy::disallowed_methods, + reason = "this crate is the one place reqwest clients are built" +)] + +mod client; mod config; mod error; pub mod media; @@ -9,7 +16,8 @@ mod settings; mod tls; pub mod transport; -pub use config::{HttpClientConfig, Resolution, Verify}; +pub use client::Client; +pub use config::{ClientIdentity, HttpClientConfig, Resolution, Verify}; pub use error::{Error, TlsSource}; pub use pool::{ClientVariant, HttpClientPool}; pub use proxy::EnvironmentProxies; diff --git a/litellm-rust/crates/http/src/media.rs b/litellm-rust/crates/http/src/media.rs index 1b9159973ef..1dac68305b0 100644 --- a/litellm-rust/crates/http/src/media.rs +++ b/litellm-rust/crates/http/src/media.rs @@ -12,7 +12,7 @@ use reqwest::{ dns::{Addrs, Name, Resolve, Resolving}, }; -use crate::{ClientVariant, HttpClientConfig, HttpClientPool}; +use crate::{Client, ClientVariant, HttpClientConfig, HttpClientPool}; #[derive(Debug, thiserror::Error)] pub enum Error { @@ -93,8 +93,8 @@ type ProxyMatch = Arc bool + Send + Sync>; #[derive(Clone)] pub struct MediaFetcher { - pinned: reqwest::Client, - unpinned: reqwest::Client, + pinned: Client, + unpinned: Client, uses_proxy: ProxyMatch, address_resolver: Arc, url_policy: UrlPolicy, @@ -154,7 +154,7 @@ impl MediaFetcher { } #[cfg(any(test, feature = "test-support"))] - pub fn for_test(client: reqwest::Client) -> Self { + pub fn for_test(client: Client) -> Self { Self { pinned: client.clone(), unpinned: client, @@ -230,7 +230,7 @@ impl MediaFetcher { } } - async fn client_for(&self, url: &Url) -> Result<&reqwest::Client, Error> { + async fn client_for(&self, url: &Url) -> Result<&Client, Error> { if !self.url_policy.validate { return Ok(&self.unpinned); } @@ -520,10 +520,7 @@ mod tests { b"HTTP/1.1 200 OK\r\nContent-Type: application/pdf; charset=binary\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc", ) .await; - let client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("test client builds"); + let client = Client::no_redirect_for_test(); let media = MediaFetcher::for_test(client) .fetch(url, policy(3, 0)) .await @@ -539,10 +536,7 @@ mod tests { b"HTTP/1.1 200 OK\r\nContent-Type: application/pdf\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc", ) .await; - let client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("test client builds"); + let client = Client::no_redirect_for_test(); let error = MediaFetcher::for_test(client) .fetch(url, policy(2, 0)) .await @@ -557,10 +551,7 @@ mod tests { b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n2\r\nab\r\n2\r\ncd\r\n0\r\n\r\n", ) .await; - let client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("test client builds"); + let client = Client::no_redirect_for_test(); let error = MediaFetcher::for_test(client) .fetch(url, policy(3, 0)) .await diff --git a/litellm-rust/crates/http/src/outbound.rs b/litellm-rust/crates/http/src/outbound.rs index d100bdf624b..c2cfb00d79b 100644 --- a/litellm-rust/crates/http/src/outbound.rs +++ b/litellm-rust/crates/http/src/outbound.rs @@ -107,7 +107,7 @@ impl OutboundRequest { self.timeout } - pub async fn send(self, client: &reqwest::Client) -> Result { + pub async fn send(self, client: &crate::Client) -> Result { let builder = with_headers( client.post(&self.url).body(self.body), &self.headers, diff --git a/litellm-rust/crates/http/src/pool.rs b/litellm-rust/crates/http/src/pool.rs index ee47e5dc52a..1187c34f2d7 100644 --- a/litellm-rust/crates/http/src/pool.rs +++ b/litellm-rust/crates/http/src/pool.rs @@ -6,7 +6,7 @@ use std::{ use reqwest::dns::Resolve; -use crate::{config::HttpClientConfig, error::Error, proxy::EnvironmentProxies}; +use crate::{client::Client, config::HttpClientConfig, error::Error, proxy::EnvironmentProxies}; #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] pub enum ClientVariant { @@ -48,7 +48,7 @@ impl HttpClientPool { &self, config: &HttpClientConfig, variant: ClientVariant, - ) -> Result { + ) -> Result { let effective = match variant { ClientVariant::Media => HttpClientConfig { client_certificate: None, @@ -65,7 +65,7 @@ impl HttpClientPool { if let Some(pooled) = self.lock().get(&key) && pooled.built_at.elapsed() < self.ttl { - return Ok(pooled.client.clone()); + return Ok(Client::new(pooled.client.clone())); } let client = self .apply(variant, reqwest::ClientBuilder::try_from(&key.0)?) @@ -77,7 +77,7 @@ impl HttpClientPool { built_at: Instant::now(), }, ); - Ok(client) + Ok(Client::new(client)) } fn lock(&self) -> MutexGuard<'_, Clients> { @@ -116,7 +116,7 @@ mod tests { }; use super::*; - use crate::{HttpSettings, Resolution, Verify}; + use crate::{ClientIdentity, HttpSettings, Resolution, Verify}; struct FixedResolver(SocketAddr); @@ -288,7 +288,9 @@ mod tests { fn media_variant_never_loads_the_client_certificate() { let pool = pool(); let with_identity = HttpClientConfig { - client_certificate: Some(std::env::temp_dir().join("litellm-http-absent-client.pem")), + client_certificate: Some(ClientIdentity::Pem( + std::env::temp_dir().join("litellm-http-absent-client.pem"), + )), ..config("a") }; assert!( diff --git a/litellm-rust/crates/http/src/request.rs b/litellm-rust/crates/http/src/request.rs index fcf296793a5..17f2e652a92 100644 --- a/litellm-rust/crates/http/src/request.rs +++ b/litellm-rust/crates/http/src/request.rs @@ -87,6 +87,20 @@ pub fn has_header(headers: &[(String, String)], name: &str) -> bool { .any(|(key, _)| key.eq_ignore_ascii_case(name)) } +pub fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> { + headers + .iter() + .find(|(key, _)| key.eq_ignore_ascii_case(name)) + .map(|(_, value)| value.as_str()) +} + +pub fn without_headers(headers: Vec<(String, String)>, names: &[&str]) -> Vec<(String, String)> { + headers + .into_iter() + .filter(|(key, _)| !names.iter().any(|name| key.eq_ignore_ascii_case(name))) + .collect() +} + pub fn has_bearer_auth(headers: &[(String, String)]) -> bool { headers.iter().any(|(name, value)| { if !name.eq_ignore_ascii_case("authorization") { @@ -194,6 +208,30 @@ mod tests { assert!(!has_header(&headers, "authorization")); } + #[test] + fn header_value_reads_the_first_match_in_any_case() { + let headers = vec![ + ("X-Api-Key".to_string(), "first".to_string()), + ("x-api-key".to_string(), "second".to_string()), + ]; + assert_eq!(header_value(&headers, "x-API-key"), Some("first")); + assert_eq!(header_value(&headers, "authorization"), None); + } + + #[test] + fn without_headers_drops_every_casing_of_the_named_headers_and_keeps_order() { + let headers = vec![ + ("X-Api-Key".to_string(), "k".to_string()), + ("anthropic-version".to_string(), "v".to_string()), + ("AUTHORIZATION".to_string(), "Bearer t".to_string()), + ("x-api-key".to_string(), "k2".to_string()), + ]; + assert_eq!( + without_headers(headers, &["x-api-key", "authorization"]), + vec![("anthropic-version".to_string(), "v".to_string())] + ); + } + #[test] fn auth_header_detection_is_case_insensitive() { let headers = vec![ diff --git a/litellm-rust/crates/http/src/tls.rs b/litellm-rust/crates/http/src/tls.rs index e2e6d27cd54..c58076607e4 100644 --- a/litellm-rust/crates/http/src/tls.rs +++ b/litellm-rust/crates/http/src/tls.rs @@ -8,7 +8,7 @@ use rustls::{ }; use crate::{ - config::{HttpClientConfig, Verify}, + config::{ClientIdentity, HttpClientConfig, Verify}, error::{Error, TlsSource}, }; @@ -203,11 +203,15 @@ impl TryFrom<&HttpClientConfig> for ClientConfig { }; let mut tls = match &config.client_certificate { None => verified.with_no_client_auth(), - Some(path) => { - let (chain, key) = identity(path, TlsSource::ClientIdentity)?; + Some(identity) => { + let (certificate, key) = match identity { + ClientIdentity::Pem(path) => (path, path), + ClientIdentity::Split { certificate, key } => (certificate, key), + }; + let (chain, private_key) = client_identity(certificate, key)?; verified - .with_client_auth_cert(chain, key) - .map_err(|error| invalid_pem(path, TlsSource::ClientIdentity, error))? + .with_client_auth_cert(chain, private_key) + .map_err(|error| invalid_pem(key, TlsSource::ClientIdentity, error))? } }; tls.alpn_protocols = if config.http2 { @@ -233,17 +237,18 @@ fn bundle_roots(path: &Path, source: TlsSource) -> Result Ok(store) } -fn identity( - path: &Path, - source: TlsSource, +fn client_identity( + certificate: &Path, + key: &Path, ) -> Result<(Vec>, PrivateKeyDer<'static>), Error> { - let chain = certificates(path, source)?; + let source = TlsSource::ClientIdentity; + let chain = certificates(certificate, source)?; if chain.is_empty() { - return Err(invalid_pem(path, source, "no certificates found")); + return Err(invalid_pem(certificate, source, "no certificates found")); } - let key = PrivateKeyDer::from_pem_slice(&read(path, source)?) - .map_err(|error| invalid_pem(path, source, error))?; - Ok((chain, key)) + let private_key = PrivateKeyDer::from_pem_slice(&read(key, source)?) + .map_err(|error| invalid_pem(key, source, error))?; + Ok((chain, private_key)) } fn certificates(path: &Path, source: TlsSource) -> Result>, Error> { @@ -405,7 +410,7 @@ mod tests { ) .unwrap(); let result = ClientConfig::try_from(&HttpClientConfig { - client_certificate: Some(path.clone()), + client_certificate: Some(ClientIdentity::Pem(path.clone())), ..config(HttpSettings::default()) }) .map(drop); @@ -419,4 +424,29 @@ mod tests { }) if reported == path )); } + + #[test] + fn split_client_identity_reads_the_key_from_its_own_file() { + let identity = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); + let directory = tempfile::tempdir().unwrap(); + let certificate = directory.path().join("client.crt"); + let key = directory.path().join("client.key"); + std::fs::write(&certificate, identity.cert.pem()).unwrap(); + std::fs::write(&key, identity.signing_key.serialize_pem()).unwrap(); + + let split = ClientConfig::try_from(&HttpClientConfig { + client_certificate: Some(ClientIdentity::Split { + certificate: certificate.clone(), + key, + }), + ..config(HttpSettings::default()) + }); + let combined = ClientConfig::try_from(&HttpClientConfig { + client_certificate: Some(ClientIdentity::Pem(certificate)), + ..config(HttpSettings::default()) + }); + + assert!(split.unwrap().client_auth_cert_resolver.has_certs()); + assert!(combined.is_err()); + } } diff --git a/litellm-rust/crates/litellm/Cargo.toml b/litellm-rust/crates/litellm/Cargo.toml new file mode 100644 index 00000000000..41009ceaff1 --- /dev/null +++ b/litellm-rust/crates/litellm/Cargo.toml @@ -0,0 +1,4 @@ +[package] +name = "litellm" +version = "0.0.1" +edition.workspace = true diff --git a/litellm-rust/crates/litellm/src/lib.rs b/litellm-rust/crates/litellm/src/lib.rs new file mode 100644 index 00000000000..ae6daac0100 --- /dev/null +++ b/litellm-rust/crates/litellm/src/lib.rs @@ -0,0 +1,2 @@ +//! Before publishing this crate, add a registry `version` beside each internal `path` dependency in the workspace manifest. +//! https://crates.io/crates/litellm diff --git a/litellm-rust/crates/llms/Cargo.toml b/litellm-rust/crates/llms/Cargo.toml index ed15d9f7cdb..beff99bc73a 100644 --- a/litellm-rust/crates/llms/Cargo.toml +++ b/litellm-rust/crates/llms/Cargo.toml @@ -11,7 +11,7 @@ test-support = ["litellm-http/test-support"] [dependencies] litellm-types.workspace = true litellm-core-utils.workspace = true -litellm-auth.workspace = true +litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] } litellm-auth-aws.workspace = true litellm-auth-azure.workspace = true litellm-auth-gcp.workspace = true @@ -19,6 +19,7 @@ litellm-host.workspace = true litellm-framing.workspace = true litellm-http.workspace = true litellm-secrets.workspace = true +litellm-python-compat.workspace = true base64.workspace = true bytes.workspace = true data-url = "0.3.2" @@ -35,6 +36,7 @@ tokio = { workspace = true, features = ["sync"] } url.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } aws-smithy-eventstream = "=0.61.4" aws-smithy-types = "1.6.1" rstest.workspace = true diff --git a/litellm-rust/crates/llms/src/anthropic/batches/AGENTS.md b/litellm-rust/crates/llms/src/anthropic/batches/AGENTS.md new file mode 100644 index 00000000000..c3a3c234492 --- /dev/null +++ b/litellm-rust/crates/llms/src/anthropic/batches/AGENTS.md @@ -0,0 +1 @@ +- https://platform.claude.com/docs/en/api/http/beta/messages/batches/create diff --git a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs index 94e4dc7838a..d1a8bdff55a 100644 --- a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs @@ -4,10 +4,7 @@ use serde_json::Value; use time::OffsetDateTime; use url::Url; -use crate::{ - anthropic::experimental_pass_through::messages::transformation::resolve_anthropic_api_base, - base_llm::chat::transformation::Error, -}; +use crate::{Error, anthropic::common_utils::resolve_anthropic_api_base}; const BATCHES_PATH_SUFFIX: &str = "/v1/messages/batches"; diff --git a/litellm-rust/crates/llms/src/anthropic/chat/handler.rs b/litellm-rust/crates/llms/src/anthropic/chat/handler.rs index a80cfbf28bd..f258656494a 100644 --- a/litellm-rust/crates/llms/src/anthropic/chat/handler.rs +++ b/litellm-rust/crates/llms/src/anthropic/chat/handler.rs @@ -7,11 +7,12 @@ use litellm_types::{ use serde_json::Value; use crate::{ - anthropic::experimental_pass_through::messages::streaming_iterator::{ + Error, + anthropic::messages::streaming_iterator::{ AnthropicContentBlock, AnthropicContentBlockDelta, AnthropicMessagesStreamEvent, AnthropicStreamUsage, }, - base_llm::{base_model_iterator::StreamTransformer, chat::transformation::Error}, + base_llm::{base_model_iterator::StreamTransformer, chat::streaming::StreamShape}, }; #[derive(Clone, Copy, Debug, Eq, PartialEq)] @@ -38,7 +39,7 @@ pub struct AnthropicContentBlockDeltaEvent { pub delta: AnthropicContentBlockDelta, } -pub struct AnthropicChatCompletionsStreamTransformer { +pub struct ModelResponseIterator { pub content_blocks: Vec, pub tool_index: i64, pub json_mode: bool, @@ -61,12 +62,8 @@ pub struct AnthropicChatCompletionsStreamTransformer { pub container_id: Option, } -impl AnthropicChatCompletionsStreamTransformer { - pub fn new( - _json_mode: bool, - _speed: Option, - _tool_name_reverse_map: HashMap, - ) -> Self { +impl ModelResponseIterator { + pub fn new(_shape: StreamShape) -> Self { todo!() } @@ -150,7 +147,7 @@ impl AnthropicChatCompletionsStreamTransformer { } } -impl StreamTransformer for AnthropicChatCompletionsStreamTransformer { +impl StreamTransformer for ModelResponseIterator { type Input = AnthropicMessagesStreamEvent; type Output = ChatCompletionChunk; type Error = Error; diff --git a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs index fd86c5ca25a..ba77a6ed790 100644 --- a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs @@ -1,3 +1,4 @@ +use litellm_auth::SecretValue; use litellm_core_utils::{ core_helpers::{finish_reason_for, unix_now, usage_from_parts}, prompt_templates::factory::{Conversation, build_conversation}, @@ -9,15 +10,24 @@ use litellm_types::{ use serde_json::{Map, Value, json}; use crate::{ + Error, anthropic::{ - ANTHROPIC_OAUTH_TOKEN_PREFIX, - experimental_pass_through::messages::transformation::{ - complete_anthropic_url, resolve_anthropic_api_key, + chat::handler::ModelResponseIterator, + common_utils::{ + API_KEY_PLACEMENT, complete_anthropic_url, forwarded_oauth_bearer, + resolve_anthropic_api_key, }, }, - base_llm::chat::transformation::{ - BaseConfig, Error, ProviderChatRequestData, ProviderChatResponseData, RequestAuth, - Unsupported, unsupported_message, unsupported_param, + base_llm::{ + anthropic_messages::streaming::anthropic_sse_event_stream, + auth::AuthScheme, + chat::{ + streaming::{ChatStream, StreamShape}, + transformation::{ + BaseConfig, Headers, ProviderChatRequestData, ProviderChatResponseData, + Unsupported, ValidatedEnvironment, unsupported_message, unsupported_param, + }, + }, }, }; @@ -65,6 +75,7 @@ impl BaseConfig for AnthropicConfig { ) -> Result { Ok(ProviderChatRequestData { body: anthropic_body(model, &build_conversation(&messages), optional_params), + stream_shape: StreamShape::default(), }) } @@ -131,17 +142,35 @@ impl BaseConfig for AnthropicConfig { }) } - fn auth( + /// A forwarded OAuth bearer is the whole credential: Python pops `x-api-key` for it, + /// so the resolved key is not applied over it. Any other forwarded header loses to + /// the deployment's key, which Python writes last. + fn validate_environment( &self, + headers: Headers, api_key: Option<&str>, _model: &str, _optional_params: &Map, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - Ok(RequestAuth::Header { - name: "x-api-key", - value: resolve_anthropic_api_key(api_key, env_lookup)?, - }) + ) -> Result { + if forwarded_oauth_bearer(&headers).is_some() { + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Forwarded, + }); + } + let auth = AuthScheme::Credential { + placement: API_KEY_PLACEMENT, + secret: SecretValue::new(resolve_anthropic_api_key(api_key, env_lookup)?), + }; + Ok(ValidatedEnvironment { headers, auth }) + } + + fn model_response_iterator(&self, shape: StreamShape) -> Option { + Some(ChatStream::new( + anthropic_sse_event_stream, + ModelResponseIterator::new(shape), + )) } fn default_headers(&self) -> &'static [(&'static str, &'static str)] { @@ -156,15 +185,6 @@ impl BaseConfig for AnthropicConfig { /// the resolved key must not be applied over the top. Any other forwarded /// `authorization` is unrelated to this header and does not defer, which is /// also what Python does: it sends the deployment's `x-api-key` alongside. - fn defers_to_forwarded_auth(&self, headers: &[(String, String)]) -> bool { - headers.iter().any(|(name, value)| { - name.eq_ignore_ascii_case("authorization") - && value - .strip_prefix("Bearer ") - .is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) - }) - } - fn unsupported_reason( &self, messages: &[ChatMessage], diff --git a/litellm-rust/crates/llms/src/anthropic/common_utils.rs b/litellm-rust/crates/llms/src/anthropic/common_utils.rs index a2234e0df03..d8d24ec0402 100644 --- a/litellm-rust/crates/llms/src/anthropic/common_utils.rs +++ b/litellm-rust/crates/llms/src/anthropic/common_utils.rs @@ -1,63 +1,33 @@ -use litellm_types::llms::anthropic_messages::anthropic_request::{ - AnthropicMessage, ContentBlock, MessageContent, +use litellm_auth::{CredentialPlacement, SecretValue}; +use litellm_http::request::{has_header, header_value, without_headers}; +use litellm_types::llms::{ + anthropic::{AnthropicBeta, BetaSet}, + anthropic_messages::anthropic_request::{ + AnthropicMessage, AnthropicTool, ContentBlock, EffortLevel, MessageContent, + }, }; +use litellm_types::recognized::Recognized; use serde::{Deserialize, Serialize}; use serde_json::Value; -use crate::anthropic::ANTHROPIC_OAUTH_TOKEN_PREFIX; +use crate::{ + anthropic::ANTHROPIC_OAUTH_TOKEN_PREFIX, + base_llm::auth::{AuthScheme, Headers}, +}; -pub const ANTHROPIC_OAUTH_BETA_HEADER: &str = "oauth-2025-04-20"; -pub const ANTHROPIC_ADVISOR_TOOL_TYPE: &str = "advisor_20260301"; -pub const ANTHROPIC_TOOL_SEARCH_TOOL_TYPES: [&str; 2] = [ - "tool_search_tool_regex_20251119", - "tool_search_tool_bm25_20251119", -]; +pub const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY"; +pub const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN"; pub const ENCRYPTED_REASONING_SIGNATURE_PREFIX: &str = "litellm_encrypted_reasoning:"; const THOUGHT_SIGNATURE_SEPARATOR: &str = "__thought__"; - -pub mod beta { - pub const CONTEXT_MANAGEMENT_2025_06_27: &str = "context-management-2025-06-27"; - pub const COMPACT_2026_01_12: &str = "compact-2026-01-12"; - pub const COMPACT_2026_09_04: &str = "compact-2026-09-04"; - pub const STRUCTURED_OUTPUT: &str = "structured-outputs-2025-11-13"; - pub const ADVANCED_TOOL_USE_2025_11_20: &str = "advanced-tool-use-2025-11-20"; - pub const FAST_MODE_2026_02_01: &str = "fast-mode-2026-02-01"; - pub const ADVISOR_TOOL_2026_03_01: &str = "advisor-tool-2026-03-01"; - pub const PER_TURN_CONTROL_2026_07_01: &str = "per-turn-control-2026-07-01"; -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -pub enum EffortLevel { - Low, - Medium, - High, - Xhigh, - Max, -} - -impl EffortLevel { - pub fn as_str(self) -> &'static str { - match self { - Self::Low => "low", - Self::Medium => "medium", - Self::High => "high", - Self::Xhigh => "xhigh", - Self::Max => "max", - } - } - - pub fn parse(value: &str) -> Option { - match value { - "low" => Some(Self::Low), - "medium" => Some(Self::Medium), - "high" => Some(Self::High), - "xhigh" => Some(Self::Xhigh), - "max" => Some(Self::Max), - _ => None, - } - } -} +const BETA_HEADER: &str = "anthropic-beta"; +pub const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE"; +pub const ANTHROPIC_BASE_URL_ENV: &str = "ANTHROPIC_BASE_URL"; +pub const DEFAULT_ANTHROPIC_API_BASE: &str = "https://api.anthropic.com"; +pub const MESSAGES_PATH_SUFFIX: &str = "/v1/messages"; +pub const API_KEY_PLACEMENT: CredentialPlacement = CredentialPlacement::Header("x-api-key"); +const API_KEY_HEADER: &str = API_KEY_PLACEMENT.header_name(); +const AUTHORIZATION: &str = CredentialPlacement::Bearer.header_name(); +const DIRECT_BROWSER_ACCESS_HEADER: &str = "anthropic-dangerous-direct-browser-access"; #[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] pub struct SupportedEffortTiers { @@ -135,55 +105,216 @@ impl AnthropicModelCapabilities { self.supports_output_config || self.effort_tiers.any() } - pub fn effort_level_rejection(&self, effort: &str, model: &str) -> Option { - match effort { - "max" if !(self.supports_adaptive_thinking || self.effort_tiers.max) => Some(format!( - "effort='max' is not supported by this model. Got model: {model}" - )), - "xhigh" if !self.effort_tiers.xhigh => Some(format!( - "effort='xhigh' is not supported by this model. Got model: {model}" - )), - _ => None, + pub fn accepts_effort(&self, level: EffortLevel) -> bool { + match level { + EffortLevel::Max => self.supports_adaptive_thinking || self.effort_tiers.max, + EffortLevel::Xhigh => self.effort_tiers.xhigh, + EffortLevel::Low | EffortLevel::Medium | EffortLevel::High => true, } } } -pub fn is_anthropic_oauth_key(value: &str) -> bool { - value - .strip_prefix("Bearer ") - .unwrap_or(value) - .starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX) +pub fn non_empty(value: Option<&str>) -> Option<&str> { + value.map(str::trim).filter(|value| !value.is_empty()) } -pub fn split_beta_values(header: Option<&str>) -> impl Iterator + '_ { - header - .into_iter() - .flat_map(|value| value.split(',')) - .map(str::trim) - .filter(|piece| !piece.is_empty()) +pub fn non_empty_env(env_lookup: &dyn Fn(&str) -> Option, name: &str) -> Option { + env_lookup(name).filter(|value| !value.trim().is_empty()) +} + +/// An Anthropic OAuth access token, which authenticates as a bearer instead of an `x-api-key`. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct OauthToken<'a>(&'a str); + +impl<'a> OauthToken<'a> { + /// The raw token, as a caller passes it in `api_key`. + pub fn parse(value: &'a str) -> Option { + value + .starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX) + .then_some(Self(value)) + } + + /// A configured key, which users paste either raw or already prefixed with `Bearer `. + pub fn parse_key(value: &'a str) -> Option { + Self::parse(value.strip_prefix("Bearer ").unwrap_or(value)) + } + + pub fn as_str(self) -> &'a str { + self.0 + } + + pub fn into_auth(self) -> AuthScheme { + AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret: SecretValue::new(self.0), + } + } +} + +/// Python's `AnthropicModelInfo.get_api_key`: the param, else `ANTHROPIC_API_KEY`. +pub fn get_api_key( + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> Option { + non_empty(api_key) .map(str::to_string) + .or_else(|| non_empty_env(env_lookup, ANTHROPIC_API_KEY_ENV)) } -pub fn join_beta_values(values: impl IntoIterator) -> String { - let mut values: Vec = values.into_iter().collect(); - values.sort(); - values.dedup(); - values.join(",") +pub fn get_auth_token(env_lookup: &dyn Fn(&str) -> Option) -> Option { + non_empty_env(env_lookup, ANTHROPIC_AUTH_TOKEN_ENV) } -pub fn is_tool_search_used(tools: Option<&[Value]>) -> bool { - tools.into_iter().flatten().any(|tool| { - tool.get("type") - .and_then(Value::as_str) - .is_some_and(|tool_type| ANTHROPIC_TOOL_SEARCH_TOOL_TYPES.contains(&tool_type)) +/// Python's `AnthropicModelInfo.get_auth_header`, naming the credential instead of building +/// the header: the key goes in `x-api-key` unless it is an OAuth token, and without a key +/// `ANTHROPIC_AUTH_TOKEN` is sent as a bearer. +pub fn get_auth_header( + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> Option { + if let Some(key) = get_api_key(api_key, env_lookup) { + return Some(match OauthToken::parse_key(&key) { + Some(token) => token.into_auth(), + None => AuthScheme::Credential { + placement: API_KEY_PLACEMENT, + secret: SecretValue::new(key), + }, + }); + } + get_auth_token(env_lookup).map(|token| AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret: SecretValue::new(token), }) } -pub fn has_advisor_tool(tools: Option<&[Value]>) -> bool { +pub fn resolve_anthropic_api_key( + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> Result { + get_api_key(api_key, env_lookup).ok_or(litellm_auth::Error::MissingApiKey { + provider: "Anthropic", + environment_variable: ANTHROPIC_API_KEY_ENV, + }) +} + +/// Whether the caller already forwarded an Anthropic credential, in either header. +pub fn has_anthropic_credential(headers: &[(String, String)]) -> bool { + has_header(headers, API_KEY_HEADER) || has_header(headers, AUTHORIZATION) +} + +pub fn resolve_anthropic_api_base( + api_base: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> String { + non_empty(api_base) + .map(str::to_string) + .or_else(|| non_empty_env(env_lookup, ANTHROPIC_API_BASE_ENV)) + .or_else(|| non_empty_env(env_lookup, ANTHROPIC_BASE_URL_ENV)) + .unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string()) +} + +pub fn complete_anthropic_url( + api_base: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> String { + let api_base = resolve_anthropic_api_base(api_base, env_lookup); + + let api_base = api_base.trim_end_matches('/'); + if api_base.ends_with(MESSAGES_PATH_SUFFIX) { + return api_base.to_string(); + } + format!("{api_base}{MESSAGES_PATH_SUFFIX}") +} + +pub fn existing_betas(headers: &[(String, String)]) -> BetaSet { + headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case(BETA_HEADER)) + .flat_map(|(_, value)| { + value + .parse::() + .unwrap_or_else(|never| match never {}) + }) + .collect() +} + +/// Python's `_merge_beta_headers`, over every casing of the header at once: the union of what +/// the caller sent and `added` replaces the header, sorted and deduplicated. Headers without +/// any beta value stay as they are. +pub fn merge_beta_headers(headers: Headers, added: BetaSet) -> Headers { + let merged = existing_betas(&headers).union(added); + if merged.is_empty() { + return headers; + } + without_headers(headers, &[BETA_HEADER]) + .into_iter() + .chain([(BETA_HEADER.to_string(), merged.to_string())]) + .collect() +} + +/// The outcome of Python's `optionally_handle_anthropic_oauth`. +#[derive(Clone, Debug, PartialEq)] +pub enum OauthHandling { + /// An OAuth token is the whole credential. The headers carry its companions and no + /// longer any `x-api-key` or `authorization`, so the bearer is applied on top. + Bearer { + headers: Headers, + token: SecretValue, + }, + Untouched(Headers), +} + +/// The OAuth token a caller forwarded as `Authorization: Bearer sk-ant-oat…`. +pub fn forwarded_oauth_bearer(headers: &[(String, String)]) -> Option> { + header_value(headers, AUTHORIZATION) + .and_then(|value| value.strip_prefix("Bearer ")) + .and_then(OauthToken::parse) +} + +fn with_oauth_companions(headers: Headers, dropped: &[&str]) -> Headers { + merge_beta_headers( + without_headers(headers, dropped), + BetaSet::from_iter([AnthropicBeta::Oauth20250420]), + ) + .into_iter() + .chain([(DIRECT_BROWSER_ACCESS_HEADER.to_string(), "true".to_string())]) + .collect() +} + +pub fn optionally_handle_anthropic_oauth(headers: Headers, api_key: Option<&str>) -> OauthHandling { + if let Some(token) = + forwarded_oauth_bearer(&headers).map(|token| SecretValue::new(token.as_str())) + { + return OauthHandling::Bearer { + headers: with_oauth_companions(headers, &[API_KEY_HEADER, AUTHORIZATION]), + token, + }; + } + if let Some(token) = api_key.and_then(OauthToken::parse) { + return OauthHandling::Bearer { + headers: with_oauth_companions(headers, &[API_KEY_HEADER]), + token: SecretValue::new(token.as_str()), + }; + } + OauthHandling::Untouched(headers) +} + +pub fn is_tool_search_used(tools: Option<&[Recognized]>) -> bool { + tools.into_iter().flatten().any(|tool| { + matches!( + tool, + Recognized::Known( + AnthropicTool::ToolSearchRegex { .. } | AnthropicTool::ToolSearchBm25 { .. } + ) + ) + }) +} + +pub fn has_advisor_tool(tools: Option<&[Recognized]>) -> bool { tools .into_iter() .flatten() - .any(|tool| tool.get("type").and_then(Value::as_str) == Some(ANTHROPIC_ADVISOR_TOOL_TYPE)) + .any(|tool| matches!(tool, Recognized::Known(AnthropicTool::Advisor { .. }))) } pub fn requires_native_compaction_beta( @@ -558,8 +689,97 @@ mod tests { serde_json::from_value(messages).unwrap() } - fn tools(value: Option) -> Option> { - value.map(|tools| tools.as_array().unwrap().clone()) + fn tools(value: Option) -> Option>> { + value.map(|tools| serde_json::from_value(tools).unwrap()) + } + + fn headers(pairs: &[(&str, &str)]) -> Headers { + pairs + .iter() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect() + } + + fn betas(values: &[&str]) -> BetaSet { + values.join(",").parse().unwrap() + } + + fn env(vars: &'static [(&'static str, &'static str)]) -> impl Fn(&str) -> Option { + move |name| { + vars.iter() + .find(|(key, _)| *key == name) + .map(|(_, value)| value.to_string()) + } + } + + const BOTH_BASE_ENVS: &[(&str, &str)] = &[ + (ANTHROPIC_API_BASE_ENV, "https://api-base.example.com"), + (ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com"), + ]; + + #[rstest] + #[case::public_endpoint_by_default(None, &[], "https://api.anthropic.com")] + #[case::explicit_api_base_beats_env( + Some("https://explicit.example.com"), + BOTH_BASE_ENVS, + "https://explicit.example.com" + )] + #[case::explicit_api_base_is_trimmed( + Some(" https://explicit.example.com "), + &[], + "https://explicit.example.com" + )] + #[case::blank_api_base_falls_back_to_env( + Some(" "), + BOTH_BASE_ENVS, + "https://api-base.example.com" + )] + #[case::api_base_env_beats_base_url_env(None, BOTH_BASE_ENVS, "https://api-base.example.com")] + #[case::base_url_env_without_api_base_env( + None, + &[(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")], + "https://base-url.example.com" + )] + #[case::blank_api_base_env_falls_back_to_base_url_env( + None, + &[(ANTHROPIC_API_BASE_ENV, " \t "), (ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")], + "https://base-url.example.com" + )] + #[case::blank_envs_fall_back_to_public_endpoint( + None, + &[(ANTHROPIC_API_BASE_ENV, ""), (ANTHROPIC_BASE_URL_ENV, " ")], + "https://api.anthropic.com" + )] + fn api_base_resolution( + #[case] api_base: Option<&str>, + #[case] vars: &'static [(&'static str, &'static str)], + #[case] expected: &str, + ) { + assert_eq!(resolve_anthropic_api_base(api_base, &env(vars)), expected); + } + + #[rstest] + #[case::forwarded_api_key(&[("X-Api-Key", "k")], true)] + #[case::forwarded_bearer(&[("Authorization", "Bearer t")], true)] + #[case::nothing_forwarded(&[("anthropic-version", "2023-06-01")], false)] + fn forwarded_credential_is_detected_in_either_header( + #[case] forwarded: &[(&str, &str)], + #[case] expected: bool, + ) { + let headers: Headers = forwarded + .iter() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect(); + assert_eq!(has_anthropic_credential(&headers), expected); + } + + fn credential(auth: Option) -> Option<(&'static str, String)> { + auth.map(|auth| match auth { + AuthScheme::Credential { placement, secret } => { + (placement.header_name(), secret.expose().to_string()) + } + other => panic!("expected a credential, got {other:?}"), + }) } fn tagged(encrypted: &str) -> String { @@ -1232,50 +1452,268 @@ mod tests { assert_eq!(twice, once); } + const OAUTH_TOKEN: &str = "sk-ant-oat01-token"; + const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token"; + const REGULAR_KEY: &str = "sk-ant-api03-regular"; + const OAUTH_BETA: &str = "oauth-2025-04-20"; + const BROWSER_ACCESS: (&str, &str) = ("anthropic-dangerous-direct-browser-access", "true"); + #[rstest] - #[case::no_existing_header(None, "b", "b")] - #[case::empty_existing_header(Some(""), "b", "b")] - #[case::whitespace_existing_header(Some(" "), "b", "b")] - #[case::sorted_after_merge(Some("c,a"), "b", "a,b,c")] - #[case::already_present(Some("a,b"), "a", "a,b")] - #[case::trimmed_and_deduplicated(Some("b, a ,b"), "c", "a,b,c")] - #[case::blank_pieces_skipped(Some("a,,b"), "c", "a,b,c")] - fn beta_values_merge_sorted_and_deduplicated( - #[case] existing: Option<&str>, - #[case] new_beta: &str, - #[case] expected: &str, + #[case::no_beta_header(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])] + #[case::blank_beta_header(&[("Anthropic-Beta", " , "), ("x-api-key", "k")], &[], &[("Anthropic-Beta", " , "), ("x-api-key", "k")])] + #[case::added_to_no_header(&[("x-api-key", "k")], &["b"], &[("x-api-key", "k"), ("anthropic-beta", "b")])] + #[case::added_to_blank_header(&[("anthropic-beta", " ")], &["b"], &[("anthropic-beta", "b")])] + #[case::sorted_after_merge(&[("anthropic-beta", "c,a")], &["b"], &[("anthropic-beta", "a,b,c")])] + #[case::already_present(&[("anthropic-beta", "a,b")], &["a"], &[("anthropic-beta", "a,b")])] + #[case::existing_normalized_without_additions( + &[("Anthropic-Beta", "b, a ,b"), ("x-api-key", "k")], + &[], + &[("x-api-key", "k"), ("anthropic-beta", "a,b")] + )] + #[case::every_casing_unioned_into_one_lowercase_header( + &[("anthropic-beta", "a"), ("ANTHROPIC-BETA", "c"), ("x-api-key", "k")], + &["b"], + &[("x-api-key", "k"), ("anthropic-beta", "a,b,c")] + )] + fn merge_beta_headers_replaces_the_header_with_the_sorted_union( + #[case] input: &[(&str, &str)], + #[case] added: &[&str], + #[case] expected: &[(&str, &str)], ) { assert_eq!( - join_beta_values(split_beta_values(existing).chain([new_beta.to_string()])), + merge_beta_headers(headers(input), betas(added)), + headers(expected) + ); + } + + #[rstest] + #[case::raw_token(OAUTH_TOKEN, Some(OAUTH_TOKEN))] + #[case::bare_prefix(ANTHROPIC_OAUTH_TOKEN_PREFIX, Some(ANTHROPIC_OAUTH_TOKEN_PREFIX))] + #[case::bearer_token(OAUTH_BEARER, None)] + #[case::api_key(REGULAR_KEY, None)] + #[case::empty("", None)] + #[case::uppercase_prefix("sk-ant-OAT01-abc123", None)] + #[case::prefix_not_at_start(" sk-ant-oat01-abc123", None)] + fn oauth_token_parses_only_the_raw_token(#[case] value: &str, #[case] expected: Option<&str>) { + assert_eq!(OauthToken::parse(value).map(OauthToken::as_str), expected); + } + + #[rstest] + #[case::raw_token(OAUTH_TOKEN, Some(OAUTH_TOKEN))] + #[case::bearer_token(OAUTH_BEARER, Some(OAUTH_TOKEN))] + #[case::api_key(REGULAR_KEY, None)] + #[case::bearer_api_key("Bearer sk-ant-api01-abc123", None)] + #[case::empty("", None)] + #[case::shouting_prefix("SK-ANT-OAT01-abc123", None)] + #[case::lowercase_bearer("bearer sk-ant-oat01-abc123", None)] + #[case::bearer_stripped_once("Bearer Bearer sk-ant-oat01-abc123", None)] + fn oauth_key_parses_the_token_behind_an_optional_bearer( + #[case] value: &str, + #[case] expected: Option<&str>, + ) { + assert_eq!( + OauthToken::parse_key(value).map(OauthToken::as_str), expected ); } #[rstest] - #[case::raw_token("sk-ant-oat01-abc123", true)] - #[case::bearer_token("Bearer sk-ant-oat02-xyz789", true)] - #[case::bare_prefix(ANTHROPIC_OAUTH_TOKEN_PREFIX, true)] - #[case::api_key("sk-ant-api01-abc123", false)] - #[case::bearer_api_key("Bearer sk-ant-api01-abc123", false)] - #[case::empty("", false)] - #[case::uppercase_prefix("sk-ant-OAT01-abc123", false)] - #[case::shouting_prefix("SK-ANT-OAT01-abc123", false)] - #[case::lowercase_bearer("bearer sk-ant-oat01-abc123", false)] - #[case::bearer_stripped_once("Bearer Bearer sk-ant-oat01-abc123", false)] - #[case::prefix_not_at_start(" sk-ant-oat01-abc123", false)] - fn anthropic_oauth_key_detection(#[case] value: &str, #[case] expected: bool) { - assert_eq!(is_anthropic_oauth_key(value), expected); + #[case::bearer(&[("authorization", OAUTH_BEARER)], Some(OAUTH_TOKEN))] + #[case::uppercase_header(&[("AUTHORIZATION", OAUTH_BEARER)], Some(OAUTH_TOKEN))] + #[case::non_oauth_bearer(&[("authorization", "Bearer some-proxy-token")], None)] + #[case::token_without_the_bearer_scheme(&[("authorization", OAUTH_TOKEN)], None)] + #[case::lowercase_bearer_scheme(&[("authorization", "bearer sk-ant-oat01-token")], None)] + #[case::token_in_x_api_key(&[("x-api-key", OAUTH_TOKEN)], None)] + #[case::no_headers(&[], None)] + fn forwarded_oauth_bearer_reads_the_authorization_header( + #[case] forwarded: &[(&str, &str)], + #[case] expected: Option<&str>, + ) { + assert_eq!( + forwarded_oauth_bearer(&headers(forwarded)).map(OauthToken::as_str), + expected + ); } #[rstest] - #[case::regex_tool(Some(json!([{"type": ANTHROPIC_TOOL_SEARCH_TOOL_TYPES[0], "name": "tool_search_tool_regex"}])), true)] - #[case::bm25_tool(Some(json!([{"type": ANTHROPIC_TOOL_SEARCH_TOOL_TYPES[1], "name": "tool_search_tool_bm25"}])), true)] + #[case::forwarded_bearer_drops_forwarded_and_deployment_keys( + &[("X-Api-Key", REGULAR_KEY), ("Authorization", OAUTH_BEARER)], + Some(REGULAR_KEY), + &[], + )] + #[case::forwarded_bearer_keeps_unrelated_headers_in_place( + &[("anthropic-version", "2023-06-01"), ("authorization", OAUTH_BEARER)], + None, + &[("anthropic-version", "2023-06-01")], + )] + #[case::forwarded_bearer_wins_over_an_oauth_api_key( + &[("authorization", OAUTH_BEARER)], + Some("sk-ant-oat01-deployment"), + &[], + )] + #[case::api_key_alone(&[], Some(OAUTH_TOKEN), &[])] + #[case::api_key_removes_a_forwarded_x_api_key(&[("x-api-key", OAUTH_TOKEN)], Some(OAUTH_TOKEN), &[])] + #[case::api_key_keeps_a_forwarded_non_oauth_bearer( + &[("Authorization", "Bearer some-proxy-token")], + Some(OAUTH_TOKEN), + &[("Authorization", "Bearer some-proxy-token")], + )] + fn oauth_token_is_the_whole_credential( + #[case] forwarded: &[(&str, &str)], + #[case] api_key: Option<&str>, + #[case] kept: &[(&str, &str)], + ) { + let expected = kept + .iter() + .copied() + .chain([("anthropic-beta", OAUTH_BETA), BROWSER_ACCESS]) + .collect::>(); + assert_eq!( + optionally_handle_anthropic_oauth(headers(forwarded), api_key), + OauthHandling::Bearer { + headers: headers(&expected), + token: SecretValue::new(OAUTH_TOKEN), + } + ); + } + + #[rstest] + #[case::forwarded_bearer_merges_a_differently_cased_beta_header( + &[("Anthropic-Beta", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)], + None, + )] + #[case::forwarded_bearer_dedupes_an_existing_oauth_beta( + &[("anthropic-beta", "web-search-2025-03-05, oauth-2025-04-20"), ("authorization", OAUTH_BEARER)], + None, + )] + #[case::api_key_merges_the_existing_beta_header( + &[("anthropic-beta", " web-search-2025-03-05 ,")], + Some(OAUTH_TOKEN), + )] + #[case::forwarded_bearer_unions_every_beta_header_casing( + &[("anthropic-beta", "oauth-2025-04-20"), ("ANTHROPIC-BETA", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)], + None, + )] + fn oauth_beta_merges_into_existing_betas( + #[case] forwarded: &[(&str, &str)], + #[case] api_key: Option<&str>, + ) { + assert_eq!( + optionally_handle_anthropic_oauth(headers(forwarded), api_key), + OauthHandling::Bearer { + headers: headers(&[ + ("anthropic-beta", "oauth-2025-04-20,web-search-2025-03-05"), + BROWSER_ACCESS, + ]), + token: SecretValue::new(OAUTH_TOKEN), + } + ); + } + + #[rstest] + #[case::x_api_key(&[("x-api-key", "caller-key")], Some("sk-other"))] + #[case::non_oauth_bearer(&[("Authorization", "Bearer some-proxy-token")], Some(REGULAR_KEY))] + #[case::oauth_token_without_the_bearer_scheme(&[("authorization", OAUTH_TOKEN)], None)] + #[case::bearer_prefixed_api_key(&[], Some(OAUTH_BEARER))] + #[case::nothing(&[], None)] + fn without_an_oauth_token_the_headers_are_untouched( + #[case] forwarded: &[(&str, &str)], + #[case] api_key: Option<&str>, + ) { + assert_eq!( + optionally_handle_anthropic_oauth(headers(forwarded), api_key), + OauthHandling::Untouched(headers(forwarded)) + ); + } + + #[rstest] + #[case::api_key_param(Some("sk-param"), &[], Some(("x-api-key", "sk-param")))] + #[case::api_key_param_over_env_key_and_auth_token( + Some("sk-param"), + &[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")], + Some(("x-api-key", "sk-param")), + )] + #[case::env_key_without_a_param(None, &[("ANTHROPIC_API_KEY", "sk-env")], Some(("x-api-key", "sk-env")))] + #[case::env_key_when_the_param_is_blank(Some(" "), &[("ANTHROPIC_API_KEY", "sk-env")], Some(("x-api-key", "sk-env")))] + #[case::env_key_over_auth_token( + None, + &[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")], + Some(("x-api-key", "sk-env")), + )] + #[case::auth_token_as_a_bearer( + None, + &[("ANTHROPIC_AUTH_TOKEN", "env-token")], + Some(("Authorization", "env-token")), + )] + #[case::auth_token_when_the_env_key_is_blank( + None, + &[("ANTHROPIC_API_KEY", " \t"), ("ANTHROPIC_AUTH_TOKEN", "env-token")], + Some(("Authorization", "env-token")), + )] + #[case::oauth_param_as_a_bearer(Some(OAUTH_TOKEN), &[], Some(("Authorization", OAUTH_TOKEN)))] + #[case::bearer_prefixed_oauth_env_key_as_a_bearer_once( + None, + &[("ANTHROPIC_API_KEY", OAUTH_BEARER)], + Some(("Authorization", OAUTH_TOKEN)), + )] + #[case::no_credentials(None, &[], None)] + #[case::blank_everything(Some(""), &[("ANTHROPIC_API_KEY", " "), ("ANTHROPIC_AUTH_TOKEN", " \t")], None)] + fn auth_header_prefers_the_key_then_the_auth_token( + #[case] api_key: Option<&str>, + #[case] vars: &'static [(&'static str, &'static str)], + #[case] expected: Option<(&str, &str)>, + ) { + assert_eq!( + credential(get_auth_header(api_key, &env(vars))), + expected.map(|(header, secret)| (header, secret.to_string())) + ); + } + + #[rstest] + #[case::param(Some("sk-param"), &[("ANTHROPIC_API_KEY", "sk-env")], Ok("sk-param"))] + #[case::blank_param_falls_back_to_env(Some(" "), &[("ANTHROPIC_API_KEY", "sk-env")], Ok("sk-env"))] + #[case::env_without_param(None, &[("ANTHROPIC_API_KEY", "sk-env")], Ok("sk-env"))] + #[case::blank_env_is_missing(None, &[("ANTHROPIC_API_KEY", " ")], Err(()))] + #[case::nothing_is_missing(None, &[], Err(()))] + fn api_key_resolution( + #[case] api_key: Option<&str>, + #[case] vars: &'static [(&'static str, &'static str)], + #[case] expected: Result<&str, ()>, + ) { + assert_eq!( + resolve_anthropic_api_key(api_key, &env(vars)).map_err(|error| { + assert!(matches!( + error, + litellm_auth::Error::MissingApiKey { + provider: "Anthropic", + environment_variable: "ANTHROPIC_API_KEY", + } + )); + }), + expected.map(str::to_string) + ); + } + + #[rstest] + #[case::absent(None, None)] + #[case::blank(Some(" \t "), None)] + #[case::padded(Some(" value "), Some("value"))] + fn non_empty_trims_and_drops_blank_values( + #[case] value: Option<&str>, + #[case] expected: Option<&str>, + ) { + assert_eq!(non_empty(value), expected); + } + + #[rstest] + #[case::regex_tool(Some(json!([{"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"}])), true)] + #[case::bm25_tool(Some(json!([{"type": "tool_search_tool_bm25_20251119", "name": "tool_search_tool_bm25"}])), true)] #[case::after_other_tools( - Some(json!([{"name": "get_weather", "input_schema": {}}, {"type": ANTHROPIC_TOOL_SEARCH_TOOL_TYPES[1]}])), + Some(json!([{"name": "get_weather", "input_schema": {}}, {"type": "tool_search_tool_bm25_20251119"}])), true )] #[case::function_tool(Some(json!([{"type": "function", "function": {"name": "get_weather"}}])), false)] - #[case::name_without_type(Some(json!([{"name": ANTHROPIC_TOOL_SEARCH_TOOL_TYPES[0]}])), false)] + #[case::name_without_type(Some(json!([{"name": "tool_search_tool_regex_20251119"}])), false)] #[case::empty_tools(Some(json!([])), false)] #[case::no_tools(None, false)] fn tool_search_detection(#[case] input: Option, #[case] expected: bool) { @@ -1283,8 +1721,8 @@ mod tests { } #[rstest] - #[case::advisor_tool(Some(json!([{"type": ANTHROPIC_ADVISOR_TOOL_TYPE, "name": "advisor"}])), true)] - #[case::after_other_tools(Some(json!([{"name": "f", "input_schema": {}}, {"type": ANTHROPIC_ADVISOR_TOOL_TYPE}])), true)] + #[case::advisor_tool(Some(json!([{"type": "advisor_20260301", "name": "advisor"}])), true)] + #[case::after_other_tools(Some(json!([{"name": "f", "input_schema": {}}, {"type": "advisor_20260301"}])), true)] #[case::tool_named_advisor(Some(json!([{"name": "advisor", "input_schema": {}}])), false)] #[case::other_server_tool(Some(json!([{"type": "web_search_20250305", "name": "web_search"}])), false)] #[case::empty_tools(Some(json!([])), false)] @@ -1329,34 +1767,6 @@ mod tests { ); } - #[rstest] - #[case::low(EffortLevel::Low, "low")] - #[case::medium(EffortLevel::Medium, "medium")] - #[case::high(EffortLevel::High, "high")] - #[case::xhigh(EffortLevel::Xhigh, "xhigh")] - #[case::max(EffortLevel::Max, "max")] - fn effort_level_names_agree_across_str_parse_and_serde( - #[case] level: EffortLevel, - #[case] name: &str, - ) { - assert_eq!(level.as_str(), name); - assert_eq!(EffortLevel::parse(name), Some(level)); - assert_eq!(serde_json::to_value(level).unwrap(), json!(name)); - assert_eq!( - serde_json::from_value::(json!(name)).unwrap(), - level - ); - } - - #[rstest] - #[case::unknown("ultra")] - #[case::minimal_is_not_an_output_config_level("minimal")] - #[case::uppercase("HIGH")] - #[case::empty("")] - fn effort_level_parse_rejects(#[case] value: &str) { - assert_eq!(EffortLevel::parse(value), None); - } - #[rstest] #[case::minimal_only(tiers(true, false, false, false, false, false), [false, false, false, false, false])] #[case::low_only(tiers(false, true, false, false, false, false), [true, false, false, false, false])] @@ -1450,56 +1860,55 @@ mod tests { } #[rstest] - #[case::max_on_adaptive_thinking_model(true, SupportedEffortTiers::default(), "max", None)] + #[case::max_on_adaptive_thinking_model( + true, + SupportedEffortTiers::default(), + EffortLevel::Max, + true + )] #[case::max_on_max_tier_model( false, tiers(false, false, false, false, false, true), - "max", - None + EffortLevel::Max, + true )] #[case::max_on_output_config_only_model( false, SupportedEffortTiers::default(), - "max", - Some("effort='max' is not supported by this model. Got model: claude-test") + EffortLevel::Max, + false )] #[case::max_on_xhigh_tier_model( false, tiers(false, false, false, false, true, false), - "max", - Some("effort='max' is not supported by this model. Got model: claude-test") + EffortLevel::Max, + false )] #[case::xhigh_on_xhigh_tier_model( false, tiers(false, false, false, false, true, false), - "xhigh", - None + EffortLevel::Xhigh, + true )] #[case::xhigh_on_adaptive_thinking_model( true, SupportedEffortTiers::default(), - "xhigh", - Some("effort='xhigh' is not supported by this model. Got model: claude-test") + EffortLevel::Xhigh, + false )] #[case::xhigh_on_max_tier_model( false, tiers(false, false, false, false, false, true), - "xhigh", - Some("effort='xhigh' is not supported by this model. Got model: claude-test") + EffortLevel::Xhigh, + false )] - #[case::high_on_unmapped_model(false, SupportedEffortTiers::default(), "high", None)] - #[case::low_on_unmapped_model(false, SupportedEffortTiers::default(), "low", None)] - #[case::unknown_level_is_left_to_other_validation( - false, - SupportedEffortTiers::default(), - "ultra", - None - )] - fn effort_level_rejection_cases( + #[case::high_on_unmapped_model(false, SupportedEffortTiers::default(), EffortLevel::High, true)] + #[case::low_on_unmapped_model(false, SupportedEffortTiers::default(), EffortLevel::Low, true)] + fn accepts_effort_cases( #[case] supports_adaptive_thinking: bool, #[case] effort_tiers: SupportedEffortTiers, - #[case] effort: &str, - #[case] expected: Option<&str>, + #[case] level: EffortLevel, + #[case] expected: bool, unmapped: AnthropicModelCapabilities, ) { let capabilities = AnthropicModelCapabilities { @@ -1508,12 +1917,7 @@ mod tests { effort_tiers, ..unmapped }; - assert_eq!( - capabilities - .effort_level_rejection(effort, "claude-test") - .as_deref(), - expected - ); + assert_eq!(capabilities.accepts_effort(level), expected); } #[rstest] diff --git a/litellm-rust/crates/llms/src/anthropic/count_tokens/AGENTS.md b/litellm-rust/crates/llms/src/anthropic/count_tokens/AGENTS.md new file mode 100644 index 00000000000..f08c0c6d017 --- /dev/null +++ b/litellm-rust/crates/llms/src/anthropic/count_tokens/AGENTS.md @@ -0,0 +1 @@ +- https://platform.claude.com/docs/en/api/http/messages/count_tokens diff --git a/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs b/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs index a4d8c57ca4f..9fa831b8b66 100644 --- a/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs @@ -2,7 +2,7 @@ use litellm_types::llms::anthropic_messages::anthropic_request::{AnthropicMessag use serde::{Deserialize, Serialize}; use serde_json::Value; -use crate::{anthropic::ANTHROPIC_OAUTH_TOKEN_PREFIX, base_llm::chat::transformation::Error}; +use crate::{Error, anthropic::ANTHROPIC_OAUTH_TOKEN_PREFIX}; const COUNT_TOKENS_ENDPOINT: &str = "https://api.anthropic.com/v1/messages/count_tokens"; const TOKEN_COUNTING_BETA: &str = "token-counting-2024-11-01"; diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/headers.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/headers.rs deleted file mode 100644 index 8d48d7a0f5c..00000000000 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/headers.rs +++ /dev/null @@ -1,643 +0,0 @@ -use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; -use serde_json::Value; - -use crate::{ - anthropic::{ - ANTHROPIC_OAUTH_TOKEN_PREFIX, - common_utils::{ - ANTHROPIC_OAUTH_BETA_HEADER, beta, has_advisor_tool, is_anthropic_oauth_key, - is_tool_search_used, join_beta_values, requires_native_compaction_beta, - split_beta_values, - }, - }, - base_llm::anthropic_messages::transformation::Headers, -}; - -const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY"; -const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN"; -const BETA_HEADER: &str = "anthropic-beta"; -const AUTHORIZATION: &str = "authorization"; -const API_KEY_HEADER: &str = "x-api-key"; -const DIRECT_BROWSER_ACCESS_HEADER: &str = "anthropic-dangerous-direct-browser-access"; - -fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> { - headers - .iter() - .find(|(header, _)| header.eq_ignore_ascii_case(name)) - .map(|(_, value)| value.as_str()) -} - -fn without(headers: Headers, names: &[&str]) -> Headers { - headers - .into_iter() - .filter(|(header, _)| !names.iter().any(|name| header.eq_ignore_ascii_case(name))) - .collect() -} - -fn existing_betas(headers: &[(String, String)]) -> impl Iterator + '_ { - headers - .iter() - .filter(|(header, _)| header.eq_ignore_ascii_case(BETA_HEADER)) - .flat_map(|(_, value)| split_beta_values(Some(value))) -} - -fn with_oauth_bearer(headers: Headers, bearer: String) -> Headers { - let beta = - join_beta_values(existing_betas(&headers).chain([ANTHROPIC_OAUTH_BETA_HEADER.to_string()])); - without(headers, &[API_KEY_HEADER, AUTHORIZATION, BETA_HEADER]) - .into_iter() - .chain([ - (AUTHORIZATION.to_string(), bearer), - (BETA_HEADER.to_string(), beta), - (DIRECT_BROWSER_ACCESS_HEADER.to_string(), "true".to_string()), - ]) - .collect() -} - -fn non_empty(value: Option<&str>) -> Option<&str> { - value.map(str::trim).filter(|value| !value.is_empty()) -} - -pub fn authenticate( - headers: Headers, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - if let Some(forwarded) = header_value(&headers, AUTHORIZATION) - && forwarded - .strip_prefix("Bearer ") - .is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) - { - let bearer = forwarded.to_string(); - return Ok(with_oauth_bearer(headers, bearer)); - } - if let Some(key) = api_key.filter(|key| key.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) { - return Ok(with_oauth_bearer(headers, format!("Bearer {key}"))); - } - if header_value(&headers, API_KEY_HEADER).is_some() - || header_value(&headers, AUTHORIZATION).is_some() - { - return Ok(headers); - } - let resolved_key = non_empty(api_key) - .map(str::to_string) - .or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty())); - let auth = match resolved_key { - Some(key) if is_anthropic_oauth_key(&key) => { - (AUTHORIZATION.to_string(), format!("Bearer {key}")) - } - Some(key) => (API_KEY_HEADER.to_string(), key), - None => match env_lookup(ANTHROPIC_AUTH_TOKEN_ENV).filter(|value| !value.trim().is_empty()) - { - Some(token) => (AUTHORIZATION.to_string(), format!("Bearer {token}")), - None => { - return Err(litellm_auth::Error::MissingApiKey { - provider: "Anthropic", - environment_variable: ANTHROPIC_API_KEY_ENV, - }); - } - }, - }; - Ok(headers.into_iter().chain([auth]).collect()) -} - -fn context_management_betas( - context_management: Option<&Value>, -) -> impl Iterator { - let edits = context_management - .and_then(|value| value.get("edits")) - .and_then(Value::as_array) - .map(Vec::as_slice) - .unwrap_or(&[]); - let (compact, other) = edits.iter().fold((false, false), |(compact, other), edit| { - match edit.get("type").and_then(Value::as_str) { - Some("compact_20260112") => (true, other), - _ => (compact, true), - } - }); - compact - .then_some(beta::COMPACT_2026_01_12) - .into_iter() - .chain(other.then_some(beta::CONTEXT_MANAGEMENT_2025_06_27)) -} - -fn uses_structured_output(request: &AnthropicMessagesRequest) -> bool { - request.output_format.is_some() - || request - .output_config - .as_ref() - .and_then(|config| config.get("format")) - .is_some_and(|format| !format.is_null()) -} - -fn messages_carry_output_config(request: &AnthropicMessagesRequest) -> bool { - request - .messages - .iter() - .any(|message| message.extra.contains_key("output_config")) -} - -pub fn feature_betas(request: &AnthropicMessagesRequest) -> Vec<&'static str> { - let tools = request.tools.as_deref(); - [ - requires_native_compaction_beta(request.compaction.as_ref(), &request.messages) - .then_some(beta::COMPACT_2026_09_04), - uses_structured_output(request).then_some(beta::STRUCTURED_OUTPUT), - (request.speed.as_deref() == Some("fast")).then_some(beta::FAST_MODE_2026_02_01), - messages_carry_output_config(request).then_some(beta::PER_TURN_CONTROL_2026_07_01), - has_advisor_tool(tools).then_some(beta::ADVISOR_TOOL_2026_03_01), - is_tool_search_used(tools).then_some(beta::ADVANCED_TOOL_USE_2025_11_20), - ] - .into_iter() - .flatten() - .chain(context_management_betas( - request.context_management.as_ref(), - )) - .collect() -} - -pub fn with_feature_betas(headers: Headers, request: &AnthropicMessagesRequest) -> Headers { - let existing = existing_betas(&headers).collect::>(); - let features = feature_betas(request); - if existing.is_empty() && features.is_empty() { - return headers; - } - let merged = join_beta_values( - existing - .into_iter() - .chain(features.into_iter().map(str::to_string)), - ); - without(headers, &[BETA_HEADER]) - .into_iter() - .chain([(BETA_HEADER.to_string(), merged)]) - .collect() -} - -#[cfg(test)] -mod tests { - use rstest::{fixture, rstest}; - use serde_json::json; - - use super::*; - - const OAUTH_TOKEN: &str = "sk-ant-oat01-token"; - const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token"; - const REGULAR_KEY: &str = "sk-ant-api03-regular"; - const BROWSER_ACCESS: (&str, &str) = ("anthropic-dangerous-direct-browser-access", "true"); - - type Env = &'static [(&'static str, &'static str)]; - - fn request(fields: Value) -> AnthropicMessagesRequest { - let mut body = - json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]}); - body.as_object_mut() - .unwrap() - .extend(fields.as_object().unwrap().clone()); - serde_json::from_value(body).unwrap() - } - - fn headers(pairs: &[(&str, &str)]) -> Headers { - pairs - .iter() - .map(|(name, value)| (name.to_string(), value.to_string())) - .collect() - } - - fn betas(values: &[&str]) -> String { - values.join(",") - } - - #[fixture] - fn no_env() -> Env { - &[] - } - - #[fixture] - fn full_env() -> Env { - &[ - ("ANTHROPIC_API_KEY", "sk-env"), - ("ANTHROPIC_AUTH_TOKEN", "env-token"), - ] - } - - fn authenticate_with( - forwarded: &[(&str, &str)], - api_key: Option<&str>, - env: Env, - ) -> Result { - let lookup = |name: &str| { - env.iter() - .find(|(key, _)| *key == name) - .map(|(_, value)| value.to_string()) - }; - authenticate(headers(forwarded), api_key, &lookup) - } - - #[rstest] - #[case::forwarded_bearer_drops_forwarded_and_deployment_keys( - &[("X-Api-Key", REGULAR_KEY), ("Authorization", OAUTH_BEARER)], - Some(REGULAR_KEY), - OAUTH_BEARER, - &[], - )] - #[case::forwarded_bearer_in_uppercase_authorization_header( - &[("AUTHORIZATION", OAUTH_BEARER)], - None, - OAUTH_BEARER, - &[], - )] - #[case::forwarded_bearer_keeps_unrelated_headers_in_place( - &[("anthropic-version", "2023-06-01"), ("authorization", OAUTH_BEARER)], - None, - OAUTH_BEARER, - &[("anthropic-version", "2023-06-01")], - )] - #[case::forwarded_bearer_wins_over_an_oauth_api_key( - &[("authorization", OAUTH_BEARER)], - Some("sk-ant-oat01-deployment"), - OAUTH_BEARER, - &[], - )] - #[case::api_key_authenticates_as_a_bearer(&[], Some(OAUTH_TOKEN), OAUTH_BEARER, &[])] - #[case::api_key_removes_a_forwarded_x_api_key( - &[("x-api-key", OAUTH_TOKEN)], - Some(OAUTH_TOKEN), - OAUTH_BEARER, - &[], - )] - #[case::api_key_replaces_a_forwarded_non_oauth_bearer( - &[("Authorization", "Bearer some-proxy-token")], - Some(OAUTH_TOKEN), - OAUTH_BEARER, - &[], - )] - fn oauth_token_is_the_whole_credential( - #[case] forwarded: &[(&str, &str)], - #[case] api_key: Option<&str>, - #[case] expected_bearer: &str, - #[case] kept: &[(&str, &str)], - full_env: Env, - ) { - let expected = kept - .iter() - .copied() - .chain([ - ("authorization", expected_bearer), - ("anthropic-beta", ANTHROPIC_OAUTH_BETA_HEADER), - BROWSER_ACCESS, - ]) - .collect::>(); - assert_eq!( - authenticate_with(forwarded, api_key, full_env).unwrap(), - headers(&expected) - ); - } - - #[rstest] - #[case::forwarded_bearer_merges_a_differently_cased_beta_header( - &[("Anthropic-Beta", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)], - None, - )] - #[case::forwarded_bearer_dedupes_an_existing_oauth_beta( - &[("anthropic-beta", "web-search-2025-03-05, oauth-2025-04-20"), ("authorization", OAUTH_BEARER)], - None, - )] - #[case::api_key_merges_the_existing_beta_header( - &[("anthropic-beta", " web-search-2025-03-05 ,")], - Some(OAUTH_TOKEN), - )] - #[case::forwarded_bearer_unions_every_beta_header_casing( - &[("anthropic-beta", "oauth-2025-04-20"), ("ANTHROPIC-BETA", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)], - None, - )] - fn oauth_beta_merges_into_existing_betas( - #[case] forwarded: &[(&str, &str)], - #[case] api_key: Option<&str>, - no_env: Env, - ) { - assert_eq!( - authenticate_with(forwarded, api_key, no_env).unwrap(), - headers(&[ - ("authorization", OAUTH_BEARER), - ( - "anthropic-beta", - &betas(&[ANTHROPIC_OAUTH_BETA_HEADER, "web-search-2025-03-05"]) - ), - BROWSER_ACCESS, - ]) - ); - } - - #[rstest] - #[case::x_api_key_over_the_deployment_key(&[("x-api-key", "caller-key")], Some("sk-other"))] - #[case::uppercase_x_api_key(&[("X-API-KEY", "caller-key")], None)] - #[case::non_oauth_bearer(&[("Authorization", "Bearer some-proxy-token")], None)] - #[case::non_oauth_bearer_over_a_regular_api_key( - &[("authorization", "Bearer sk-ant-api03-forwarded")], - Some(REGULAR_KEY), - )] - #[case::oauth_token_without_the_bearer_scheme(&[("authorization", OAUTH_TOKEN)], None)] - #[case::oauth_token_behind_a_lowercase_bearer_scheme( - &[("authorization", "bearer sk-ant-oat01-token")], - None, - )] - fn forwarded_auth_header_is_kept_untouched( - #[case] forwarded: &[(&str, &str)], - #[case] api_key: Option<&str>, - full_env: Env, - ) { - assert_eq!( - authenticate_with(forwarded, api_key, full_env).unwrap(), - headers(forwarded) - ); - } - - #[rstest] - #[case::api_key_param(Some("sk-param"), &[], ("x-api-key", "sk-param"))] - #[case::api_key_param_over_env_key_and_auth_token( - Some("sk-param"), - &[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")], - ("x-api-key", "sk-param"), - )] - #[case::env_key_without_a_param(None, &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))] - #[case::env_key_when_the_param_is_empty(Some(""), &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))] - #[case::env_key_when_the_param_is_whitespace( - Some(" "), - &[("ANTHROPIC_API_KEY", "sk-env")], - ("x-api-key", "sk-env"), - )] - #[case::env_key_over_auth_token( - None, - &[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")], - ("x-api-key", "sk-env"), - )] - #[case::auth_token_as_a_bearer( - None, - &[("ANTHROPIC_AUTH_TOKEN", "env-token")], - ("authorization", "Bearer env-token"), - )] - #[case::auth_token_when_the_env_key_is_whitespace( - None, - &[("ANTHROPIC_API_KEY", " \t"), ("ANTHROPIC_AUTH_TOKEN", "env-token")], - ("authorization", "Bearer env-token"), - )] - #[case::oauth_env_key_as_a_plain_bearer( - None, - &[("ANTHROPIC_API_KEY", "sk-ant-oat01-env")], - ("authorization", "Bearer sk-ant-oat01-env"), - )] - fn credential_is_resolved_after_the_existing_headers( - #[case] api_key: Option<&str>, - #[case] env: Env, - #[case] expected: (&str, &str), - ) { - let forwarded = [("anthropic-beta", "web-search-2025-03-05")]; - assert_eq!( - authenticate_with(&forwarded, api_key, env).unwrap(), - headers(&[forwarded[0], expected]) - ); - } - - #[rstest] - #[case::no_credentials(&[], None, &[])] - #[case::empty_api_key(&[], Some(""), &[])] - #[case::whitespace_only_env_values( - &[], - None, - &[("ANTHROPIC_API_KEY", " "), ("ANTHROPIC_AUTH_TOKEN", " \t")], - )] - #[case::unrelated_forwarded_headers(&[("anthropic-beta", "web-search-2025-03-05")], None, &[])] - fn missing_credentials_are_an_auth_error( - #[case] forwarded: &[(&str, &str)], - #[case] api_key: Option<&str>, - #[case] env: Env, - ) { - assert!(matches!( - authenticate_with(forwarded, api_key, env), - Err(litellm_auth::Error::MissingApiKey { - provider: "Anthropic", - environment_variable: "ANTHROPIC_API_KEY", - }) - )); - } - - #[rstest] - #[case::no_features(json!({}), &[])] - #[case::output_format(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])] - #[case::null_output_format(json!({"output_format": null}), &[])] - #[case::output_config_format( - json!({"output_config": {"format": {"type": "json_schema"}, "effort": "xhigh"}}), - &[beta::STRUCTURED_OUTPUT] - )] - #[case::null_output_config_format(json!({"output_config": {"format": null}}), &[])] - #[case::top_level_output_config_without_format(json!({"output_config": {"effort": "high"}}), &[])] - #[case::fast_speed(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])] - #[case::standard_speed(json!({"speed": "standard"}), &[])] - #[case::compaction_param(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])] - #[case::empty_compaction_param(json!({"compaction": {}}), &[beta::COMPACT_2026_09_04])] - #[case::signed_compaction_block_in_history( - json!({"messages": [ - {"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": "sig"}]}, - {"role": "user", "content": "Continue"}, - ]}), - &[beta::COMPACT_2026_09_04] - )] - #[case::unsigned_compaction_block_in_history( - json!({"messages": [ - {"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": ""}]}, - {"role": "user", "content": "Continue"}, - ]}), - &[] - )] - #[case::advisor_tool( - json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": "claude-opus-4-6"}]}), - &[beta::ADVISOR_TOOL_2026_03_01] - )] - #[case::no_tools(json!({"tools": []}), &[])] - #[case::regex_tool_search( - json!({"tools": [{"type": "tool_search_tool_regex_20251119"}]}), - &[beta::ADVANCED_TOOL_USE_2025_11_20] - )] - #[case::bm25_tool_search( - json!({"tools": [{"type": "tool_search_tool_bm25_20251119"}]}), - &[beta::ADVANCED_TOOL_USE_2025_11_20] - )] - #[case::unrelated_server_tool(json!({"tools": [{"type": "web_search_20250305", "name": "web_search"}]}), &[])] - #[case::only_compact_edits( - json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}), - &[beta::COMPACT_2026_01_12] - )] - #[case::only_other_edits( - json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919", "keep": {"type": "tool_uses", "value": 3}}]}}), - &[beta::CONTEXT_MANAGEMENT_2025_06_27] - )] - #[case::compact_and_other_edits( - json!({"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_tool_uses_20250919"}]}}), - &[beta::COMPACT_2026_01_12, beta::CONTEXT_MANAGEMENT_2025_06_27] - )] - #[case::edit_without_a_type(json!({"context_management": {"edits": [{}]}}), &[beta::CONTEXT_MANAGEMENT_2025_06_27])] - #[case::empty_edits(json!({"context_management": {"edits": []}}), &[])] - #[case::context_management_without_edits(json!({"context_management": {}}), &[])] - #[case::per_message_output_config( - json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}), - &[beta::PER_TURN_CONTROL_2026_07_01] - )] - #[case::per_message_null_output_config( - json!({"messages": [{"role": "user", "content": "hi", "output_config": null}]}), - &[beta::PER_TURN_CONTROL_2026_07_01] - )] - fn feature_betas_follow_the_request(#[case] fields: Value, #[case] expected: &[&str]) { - assert_eq!(feature_betas(&request(fields)), expected); - } - - #[rstest] - #[case::no_betas(&[("x-api-key", "k"), ("anthropic-version", "2023-06-01")], json!({}))] - #[case::blank_beta_header(&[("Anthropic-Beta", " , "), ("x-api-key", "k")], json!({}))] - fn headers_without_any_beta_value_are_untouched( - #[case] input: &[(&str, &str)], - #[case] fields: Value, - ) { - assert_eq!( - with_feature_betas(headers(input), &request(fields)), - headers(input) - ); - } - - #[rstest] - #[case::feature_beta_is_appended( - &[("x-api-key", "k")], - json!({"speed": "fast"}), - &[("x-api-key", "k"), ("anthropic-beta", beta::FAST_MODE_2026_02_01)], - )] - #[case::existing_betas_are_normalized_without_features( - &[("Anthropic-Beta", "web-search-2025-03-05, interleaved-thinking-2025-05-14 ,web-search-2025-03-05"), ("x-api-key", "k")], - json!({}), - &[("x-api-key", "k"), ("anthropic-beta", "interleaved-thinking-2025-05-14,web-search-2025-03-05")], - )] - #[case::existing_advisor_beta_is_kept_without_an_advisor_tool( - &[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)], - json!({"tools": []}), - &[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)], - )] - #[case::feature_already_sent_is_not_duplicated( - &[("anthropic-beta", beta::FAST_MODE_2026_02_01)], - json!({"speed": "fast"}), - &[("anthropic-beta", beta::FAST_MODE_2026_02_01)], - )] - fn feature_betas_merge_into_the_headers( - #[case] input: &[(&str, &str)], - #[case] fields: Value, - #[case] expected: &[(&str, &str)], - ) { - assert_eq!( - with_feature_betas(headers(input), &request(fields)), - headers(expected) - ); - } - - #[test] - fn differently_cased_beta_header_is_replaced_by_one_sorted_header() { - let merged = with_feature_betas( - headers(&[("Anthropic-Beta", "interleaved-thinking-2025-05-14")]), - &request( - json!({"messages": [{"role": "system", "content": "env", "output_config": {"effort": "low"}}]}), - ), - ); - assert_eq!( - merged, - headers(&[( - "anthropic-beta", - &betas(&[ - "interleaved-thinking-2025-05-14", - beta::PER_TURN_CONTROL_2026_07_01 - ]) - )]) - ); - } - - #[test] - fn every_beta_header_casing_is_unioned_into_one_header() { - let merged = with_feature_betas( - headers(&[ - ("anthropic-beta", "interleaved-thinking-2025-05-14"), - ("Anthropic-Beta", "web-search-2025-03-05"), - ]), - &request(json!({"speed": "fast"})), - ); - assert_eq!( - merged, - headers(&[( - "anthropic-beta", - &betas(&[ - beta::FAST_MODE_2026_02_01, - "interleaved-thinking-2025-05-14", - "web-search-2025-03-05" - ]) - )]) - ); - } - - #[test] - fn unknown_client_betas_survive_alongside_the_added_one() { - let client_betas = [ - "claude-code-20250219", - "interleaved-thinking-2025-05-14", - beta::CONTEXT_MANAGEMENT_2025_06_27, - beta::PER_TURN_CONTROL_2026_07_01, - "effort-2025-11-24", - ]; - let merged = with_feature_betas( - headers(&[("anthropic-beta", &betas(&client_betas))]), - &request( - json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}), - ), - ); - assert_eq!( - merged, - headers(&[( - "anthropic-beta", - &betas(&[ - "claude-code-20250219", - beta::CONTEXT_MANAGEMENT_2025_06_27, - "effort-2025-11-24", - "interleaved-thinking-2025-05-14", - beta::PER_TURN_CONTROL_2026_07_01, - ]) - )]) - ); - } - - #[test] - fn every_feature_merges_with_the_oauth_beta_sorted_and_last() { - let oauth_headers = authenticate_with(&[], Some(OAUTH_TOKEN), &[]).unwrap(); - let all_features = request(json!({ - "compaction": {"enabled": true}, - "output_format": {"type": "json_schema"}, - "speed": "fast", - "tools": [{"type": "advisor_20260301"}, {"type": "tool_search_tool_bm25_20251119"}], - "context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_thinking_20251015"}]}, - "messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}], - })); - assert_eq!( - with_feature_betas(oauth_headers, &all_features), - headers(&[ - ("authorization", OAUTH_BEARER), - BROWSER_ACCESS, - ( - "anthropic-beta", - &betas(&[ - beta::ADVANCED_TOOL_USE_2025_11_20, - beta::ADVISOR_TOOL_2026_03_01, - beta::COMPACT_2026_01_12, - beta::COMPACT_2026_09_04, - beta::CONTEXT_MANAGEMENT_2025_06_27, - beta::FAST_MODE_2026_02_01, - ANTHROPIC_OAUTH_BETA_HEADER, - beta::PER_TURN_CONTROL_2026_07_01, - beta::STRUCTURED_OUTPUT, - ]) - ), - ]) - ); - } -} diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs deleted file mode 100644 index 3f1b7ed9bcc..00000000000 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs +++ /dev/null @@ -1,291 +0,0 @@ -use base64::Engine; -use bytes::Buf; -use futures_util::{Stream, StreamExt}; -use litellm_framing::{ - aws_event_stream::{AwsEventStreamCodec, Message}, - frames, - sse::{SseCodec, SseEvent}, -}; -use serde::{Deserialize, Serialize}; -use serde_json::{Map, Value}; - -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("stream framing failed: {0}")] - StreamFraming(String), - #[error("Anthropic stream event is invalid: {0}")] - InvalidStreamEvent(String), - #[error("Bedrock event payload is invalid: {0}")] - InvalidBedrockPayload(String), - #[error("Bedrock event payload has invalid base64: {0}")] - InvalidBedrockBase64(String), -} - -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub struct AnthropicStreamUsage { - #[serde(default)] - pub input_tokens: u64, - #[serde(default)] - pub output_tokens: u64, - #[serde(default)] - pub cache_creation_input_tokens: u64, - #[serde(default)] - pub cache_read_input_tokens: u64, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub server_tool_use: Option, - #[serde(flatten)] - pub extra: Map, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct AnthropicStreamMessage { - pub id: String, - #[serde(rename = "type")] - pub message_type: String, - pub role: String, - pub model: String, - pub content: Vec, - pub stop_reason: Option, - pub stop_sequence: Option, - pub usage: AnthropicStreamUsage, - #[serde(flatten)] - pub extra: Map, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(tag = "type", rename_all = "snake_case")] -pub enum AnthropicContentBlockDelta { - TextDelta { - text: String, - }, - InputJsonDelta { - partial_json: String, - }, - #[serde(rename = "citations_delta")] - Citations { - citation: Value, - }, - ThinkingDelta { - thinking: String, - }, - SignatureDelta { - signature: String, - }, - CompactionDelta { - content: String, - }, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct AnthropicContentBlock { - #[serde(rename = "type")] - pub block_type: String, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub id: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub name: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub text: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub input: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub thinking: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub signature: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub data: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub caller: Option, - #[serde(flatten)] - pub extra: Map, -} - -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub struct AnthropicMessageDelta { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub stop_reason: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub stop_sequence: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub stop_details: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub container: Option, - #[serde(flatten)] - pub extra: Map, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct AnthropicStreamError { - #[serde(rename = "type")] - pub error_type: String, - pub message: String, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub details: Option, - #[serde(flatten)] - pub extra: Map, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(tag = "type", rename_all = "snake_case")] -pub enum AnthropicMessagesStreamEvent { - MessageStart { - message: AnthropicStreamMessage, - }, - ContentBlockStart { - index: u64, - content_block: AnthropicContentBlock, - }, - ContentBlockDelta { - index: u64, - delta: AnthropicContentBlockDelta, - }, - ContentBlockStop { - index: u64, - }, - MessageDelta { - delta: AnthropicMessageDelta, - #[serde(default, skip_serializing_if = "Option::is_none")] - usage: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - context_management: Option, - }, - MessageStop, - Ping, - Error { - error: AnthropicStreamError, - }, -} - -#[derive(Deserialize)] -struct BedrockChunkPayload { - bytes: String, -} - -pub fn decode_anthropic_sse_frame(event: SseEvent) -> Result { - serde_json::from_str(&event.data).map_err(|error| Error::InvalidStreamEvent(error.to_string())) -} - -pub fn decode_bedrock_anthropic_frame( - message: Message, -) -> Result { - let payload: BedrockChunkPayload = serde_json::from_slice(message.payload()) - .map_err(|error| Error::InvalidBedrockPayload(error.to_string()))?; - let event = base64::engine::general_purpose::STANDARD - .decode(payload.bytes) - .map_err(|error| Error::InvalidBedrockBase64(error.to_string()))?; - serde_json::from_slice(&event).map_err(|error| Error::InvalidStreamEvent(error.to_string())) -} - -pub fn direct_anthropic_event_stream( - input: S, -) -> impl Stream> + Send -where - S: Stream> + Send, - B: Buf + Send, - E: std::error::Error + Send + Sync + 'static, -{ - frames(input, SseCodec::default()).map(|event| { - decode_anthropic_sse_frame(event.map_err(|error| Error::StreamFraming(error.to_string()))?) - }) -} - -pub fn bedrock_anthropic_event_stream( - input: S, -) -> impl Stream> + Send -where - S: Stream> + Send, - B: Buf + Send, - E: std::error::Error + Send + Sync + 'static, -{ - frames(input, AwsEventStreamCodec).map(|message| { - decode_bedrock_anthropic_frame( - message.map_err(|error| Error::StreamFraming(error.to_string()))?, - ) - }) -} - -#[cfg(test)] -mod tests { - use std::io; - - use aws_smithy_eventstream::frame::write_message_to; - use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; - use base64::engine::general_purpose::STANDARD; - use bytes::Bytes; - use futures_util::TryStreamExt; - - use super::*; - - const TEXT_DELTA: &str = - r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hello"}}"#; - - #[tokio::test] - async fn direct_anthropic_sse_frames_into_typed_events() { - let wire = format!("event: content_block_delta\ndata: {TEXT_DELTA}\n\n"); - let events = direct_anthropic_event_stream(futures_util::stream::iter( - wire.as_bytes().chunks(3).map(Ok::<_, io::Error>), - )) - .try_collect::>() - .await - .unwrap(); - - assert_eq!( - events, - vec![AnthropicMessagesStreamEvent::ContentBlockDelta { - index: 0, - delta: AnthropicContentBlockDelta::TextDelta { - text: "hello".into(), - }, - }] - ); - } - - #[test] - fn decodes_citations_delta_events() { - let event = decode_anthropic_sse_frame(SseEvent { - event: Some("content_block_delta".into()), - data: r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"# - .into(), - id: None, - retry: None, - }) - .unwrap(); - - assert!(matches!( - event, - AnthropicMessagesStreamEvent::ContentBlockDelta { - delta: AnthropicContentBlockDelta::Citations { .. }, - .. - } - )); - } - - #[tokio::test] - async fn bedrock_aws_frames_into_the_same_typed_events() { - let payload = serde_json::json!({"bytes": STANDARD.encode(TEXT_DELTA)}); - let message = Message::new(Bytes::from(serde_json::to_vec(&payload).unwrap())).add_header( - Header::new(":event-type", HeaderValue::String("chunk".into())), - ); - let mut wire = Vec::new(); - write_message_to(&message, &mut wire).unwrap(); - - let events = bedrock_anthropic_event_stream(futures_util::stream::iter( - wire.chunks(3).map(Ok::<_, io::Error>), - )) - .try_collect::>() - .await - .unwrap(); - - assert_eq!( - events, - vec![AnthropicMessagesStreamEvent::ContentBlockDelta { - index: 0, - delta: AnthropicContentBlockDelta::TextDelta { - text: "hello".into(), - }, - }] - ); - } -} diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/thinking.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/thinking.rs deleted file mode 100644 index ffa4c8ffeb8..00000000000 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/thinking.rs +++ /dev/null @@ -1,1182 +0,0 @@ -use litellm_core_utils::settings::Lookup; -use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; -use serde_json::{Map, Value, json}; - -use crate::{ - anthropic::common_utils::AnthropicModelCapabilities, base_llm::chat::transformation::Error, -}; - -pub const ANTHROPIC_MIN_THINKING_BUDGET_TOKENS: u64 = 1024; - -const EFFORT_NAMES: &str = "'minimal', 'low', 'medium', 'high', 'xhigh', 'max', 'none'"; - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct ThinkingBudgets { - pub minimal: u64, - pub low: u64, - pub medium: u64, - pub high: u64, - pub xhigh: u64, - pub max: u64, -} - -impl Default for ThinkingBudgets { - fn default() -> Self { - Self { - minimal: 128, - low: 1024, - medium: 2048, - high: 4096, - xhigh: 8192, - max: 16384, - } - } -} - -impl ThinkingBudgets { - pub fn from_lookup(env: &impl Lookup) -> Self { - let defaults = Self::default(); - let tier = |name: &str, default: u64| { - env.parsed::(&format!("DEFAULT_REASONING_EFFORT_{name}_THINKING_BUDGET")) - .unwrap_or(default) - }; - Self { - minimal: tier("MINIMAL", defaults.minimal), - low: tier("LOW", defaults.low), - medium: tier("MEDIUM", defaults.medium), - high: tier("HIGH", defaults.high), - xhigh: tier("XHIGH", defaults.xhigh), - max: tier("MAX", defaults.max), - } - } - - fn for_effort(&self, reasoning_effort: &str) -> Option { - match reasoning_effort { - "low" => Some(self.low), - "medium" => Some(self.medium), - "high" => Some(self.high), - "xhigh" => Some(self.xhigh), - "max" => Some(self.max), - "minimal" => Some(self.minimal.max(ANTHROPIC_MIN_THINKING_BUDGET_TOKENS)), - _ => None, - } - } - - fn effort_for_budget( - &self, - budget_tokens: u64, - capabilities: &AnthropicModelCapabilities, - ) -> &'static str { - if budget_tokens >= self.xhigh && capabilities.effort_tiers.xhigh { - return "xhigh"; - } - if budget_tokens >= self.high { - return "high"; - } - if budget_tokens >= self.medium { - return "medium"; - } - "low" - } -} - -#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] -pub struct ThinkingContext { - pub capabilities: AnthropicModelCapabilities, - pub budgets: ThinkingBudgets, -} - -fn bad_request(message: String) -> Error { - Error::InvalidRequest(message) -} - -fn thinking_type(thinking: Option<&Value>) -> Option<&str> { - thinking?.get("type")?.as_str() -} - -fn output_config_effort(output_config: Option<&Value>) -> Option<&str> { - output_config?.get("effort")?.as_str() -} - -fn enabled_thinking(budget_tokens: u64) -> Value { - json!({"type": "enabled", "budget_tokens": budget_tokens}) -} - -fn map_reasoning_effort( - reasoning_effort: &str, - context: &ThinkingContext, -) -> Result, Error> { - if reasoning_effort == "none" { - return Ok(None); - } - if context.capabilities.supports_adaptive_thinking { - return Ok(Some(json!({"type": "adaptive", "display": "summarized"}))); - } - context - .budgets - .for_effort(reasoning_effort) - .map(|budget| Some(enabled_thinking(budget))) - .ok_or_else(|| { - bad_request(format!( - "Unmapped reasoning effort: '{reasoning_effort}'. Must be one of: {EFFORT_NAMES}." - )) - }) -} - -fn cap_thinking_budget_to_max_tokens(thinking: Value, max_tokens: Option) -> Option { - let (Some(max_tokens), Some(budget)) = ( - max_tokens, - thinking.get("budget_tokens").and_then(Value::as_u64), - ) else { - return Some(thinking); - }; - if max_tokens <= ANTHROPIC_MIN_THINKING_BUDGET_TOKENS { - return None; - } - if budget < max_tokens { - return Some(thinking); - } - Some(enabled_thinking(max_tokens - 1)) -} - -fn reasoning_effort_to_output_config_effort(reasoning_effort: &str) -> Option<&'static str> { - match reasoning_effort { - "low" | "minimal" => Some("low"), - "medium" => Some("medium"), - "high" => Some("high"), - "xhigh" => Some("xhigh"), - "max" => Some("max"), - _ => None, - } -} - -fn with_default_effort(output_config: Option, effort: &str) -> Value { - let mut config = match output_config { - Some(Value::Object(config)) => config, - _ => Map::new(), - }; - if !config.contains_key("effort") { - config.insert("effort".to_string(), Value::String(effort.to_string())); - } - Value::Object(config) -} - -fn translate_reasoning_effort( - request: AnthropicMessagesRequest, - context: &ThinkingContext, -) -> Result { - let Some(reasoning_effort) = request.reasoning_effort.clone() else { - return Ok(request); - }; - let request = AnthropicMessagesRequest { - reasoning_effort: None, - ..request - }; - let Some(mapped) = map_reasoning_effort(&reasoning_effort, context)? else { - return Ok(AnthropicMessagesRequest { - thinking: None, - output_config: None, - ..request - }); - }; - let Some(fitted) = cap_thinking_budget_to_max_tokens(mapped, request.max_tokens) else { - return Ok(request); - }; - let thinking = Some(request.thinking.clone().unwrap_or(fitted)); - if !context.capabilities.supports_adaptive_thinking { - return Ok(AnthropicMessagesRequest { - thinking, - ..request - }); - } - let effort = reasoning_effort_to_output_config_effort(&reasoning_effort).ok_or_else(|| { - bad_request(format!( - "Invalid reasoning_effort: '{reasoning_effort}'. Must be one of: {EFFORT_NAMES}" - )) - })?; - if let Some(rejection) = context - .capabilities - .effort_level_rejection(effort, &request.model) - { - return Err(bad_request(rejection)); - } - Ok(AnthropicMessagesRequest { - thinking, - output_config: Some(with_default_effort(request.output_config.clone(), effort)), - ..request - }) -} - -fn drop_disabled_thinking( - request: AnthropicMessagesRequest, - context: &ThinkingContext, -) -> AnthropicMessagesRequest { - if !context.capabilities.thinking_always_on - || thinking_type(request.thinking.as_ref()) != Some("disabled") - { - return request; - } - AnthropicMessagesRequest { - thinking: None, - ..request - } -} - -fn translate_legacy_thinking_for_adaptive_model( - request: AnthropicMessagesRequest, - context: &ThinkingContext, -) -> AnthropicMessagesRequest { - let capabilities = &context.capabilities; - if !capabilities.supports_adaptive_thinking - || capabilities.supports_legacy_thinking - || thinking_type(request.thinking.as_ref()) != Some("enabled") - { - return request; - } - let budget = request - .thinking - .as_ref() - .and_then(|thinking| thinking.get("budget_tokens")) - .and_then(Value::as_u64) - .unwrap_or(0); - let effort = context.budgets.effort_for_budget(budget, capabilities); - AnthropicMessagesRequest { - thinking: Some(json!({"type": "adaptive"})), - output_config: Some(with_default_effort(request.output_config.clone(), effort)), - ..request - } -} - -fn output_config_without_effort(output_config: Option) -> Option { - let Some(Value::Object(config)) = output_config else { - return output_config; - }; - if !config.contains_key("effort") { - return Some(Value::Object(config)); - } - let residual: Map = config - .into_iter() - .filter(|(key, _)| key != "effort") - .collect(); - (!residual.is_empty()).then_some(Value::Object(residual)) -} - -fn translate_adaptive_effort_for_non_adaptive_model( - request: AnthropicMessagesRequest, - context: &ThinkingContext, -) -> Result { - let capabilities = &context.capabilities; - if capabilities.supports_adaptive_thinking { - return Ok(request); - } - let effort = output_config_effort(request.output_config.as_ref()).map(str::to_string); - let adaptive_thinking = thinking_type(request.thinking.as_ref()) == Some("adaptive"); - if effort.is_none() && !adaptive_thinking { - return Ok(request); - } - let level_supported = effort.as_deref().is_none_or(|effort| { - capabilities - .effort_level_rejection(effort, &request.model) - .is_none() - }); - if capabilities.supports_effort_param() && (!adaptive_thinking || level_supported) { - return Ok(AnthropicMessagesRequest { - thinking: if adaptive_thinking { - None - } else { - request.thinking.clone() - }, - ..request - }); - } - let legacy = if capabilities.supports_reasoning { - map_reasoning_effort( - effort - .as_deref() - .filter(|effort| !effort.is_empty()) - .unwrap_or("medium"), - context, - )? - } else { - None - }; - let capped = - legacy.and_then(|thinking| cap_thinking_budget_to_max_tokens(thinking, request.max_tokens)); - Ok(AnthropicMessagesRequest { - thinking: capped, - output_config: output_config_without_effort(request.output_config.clone()), - ..request - }) -} - -fn drop_incompatible_temperature_for_thinking( - request: AnthropicMessagesRequest, - context: &ThinkingContext, -) -> AnthropicMessagesRequest { - if context.capabilities.supports_adaptive_thinking { - return request; - } - let pinned = request - .temperature - .is_some_and(|temperature| temperature != 1.0); - let thinking_enabled = thinking_type(request.thinking.as_ref()) == Some("enabled"); - let effort_enabled = output_config_effort(request.output_config.as_ref()).is_some(); - if !pinned || !(thinking_enabled || effort_enabled) { - return request; - } - AnthropicMessagesRequest { - temperature: None, - ..request - } -} - -pub fn translate_thinking( - request: AnthropicMessagesRequest, - context: &ThinkingContext, -) -> Result { - let request = translate_reasoning_effort(request, context)?; - let request = drop_disabled_thinking(request, context); - let request = translate_legacy_thinking_for_adaptive_model(request, context); - let request = translate_adaptive_effort_for_non_adaptive_model(request, context)?; - Ok(drop_incompatible_temperature_for_thinking(request, context)) -} - -#[cfg(test)] -mod tests { - use rstest::{fixture, rstest}; - - use super::*; - use crate::anthropic::common_utils::SupportedEffortTiers; - - const EFFORT_CHOICES: &str = "'minimal', 'low', 'medium', 'high', 'xhigh', 'max', 'none'"; - - fn request(fields: Value) -> AnthropicMessagesRequest { - let mut body = - json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]}); - body.as_object_mut() - .unwrap() - .extend(fields.as_object().unwrap().clone()); - serde_json::from_value(body).unwrap() - } - - fn context(capabilities: AnthropicModelCapabilities) -> ThinkingContext { - ThinkingContext { - capabilities, - budgets: ThinkingBudgets::default(), - } - } - - fn translate( - capabilities: AnthropicModelCapabilities, - fields: Value, - ) -> Result { - translate_thinking(request(fields), &context(capabilities)) - } - - fn overridden_budgets(overrides: &[(&str, &str)]) -> ThinkingBudgets { - let env = |name: &str| { - overrides - .iter() - .find(|(tier, _)| { - name == format!("DEFAULT_REASONING_EFFORT_{tier}_THINKING_BUDGET") - }) - .map(|(_, value)| value.to_string()) - }; - ThinkingBudgets::from_lookup(&env) - } - - fn claude_code_payload(effort: &str, max_tokens: u64) -> Value { - json!({"max_tokens": max_tokens, "thinking": {"type": "adaptive"}, "output_config": {"effort": effort}}) - } - - fn with_temperature(fields: Value, temperature: f64) -> Value { - let mut fields = fields; - fields - .as_object_mut() - .unwrap() - .insert("temperature".to_string(), json!(temperature)); - fields - } - - #[fixture] - fn haiku_3_5() -> AnthropicModelCapabilities { - AnthropicModelCapabilities::default() - } - - #[fixture] - fn haiku_4_5() -> AnthropicModelCapabilities { - AnthropicModelCapabilities { - supports_reasoning: true, - ..Default::default() - } - } - - #[fixture] - fn opus_4_5() -> AnthropicModelCapabilities { - AnthropicModelCapabilities { - supports_reasoning: true, - supports_output_config: true, - ..Default::default() - } - } - - #[fixture] - fn sonnet_4_6() -> AnthropicModelCapabilities { - AnthropicModelCapabilities { - supports_reasoning: true, - supports_adaptive_thinking: true, - supports_legacy_thinking: true, - supports_output_config: true, - effort_tiers: SupportedEffortTiers { - max: true, - ..Default::default() - }, - ..Default::default() - } - } - - #[fixture] - fn opus_4_7() -> AnthropicModelCapabilities { - AnthropicModelCapabilities { - supports_reasoning: true, - supports_adaptive_thinking: true, - supports_output_config: true, - effort_tiers: SupportedEffortTiers { - xhigh: true, - max: true, - ..Default::default() - }, - ..Default::default() - } - } - - #[fixture] - fn fable_5_1() -> AnthropicModelCapabilities { - AnthropicModelCapabilities { - thinking_always_on: true, - ..opus_4_7() - } - } - - #[fixture] - fn newfamily_6() -> AnthropicModelCapabilities { - AnthropicModelCapabilities { - supports_reasoning: true, - supports_adaptive_thinking: true, - ..Default::default() - } - } - - #[rstest] - #[case::minimal_maps_to_low(opus_4_7(), "minimal", "low")] - #[case::low(opus_4_7(), "low", "low")] - #[case::medium(opus_4_7(), "medium", "medium")] - #[case::high(opus_4_7(), "high", "high")] - #[case::xhigh_with_xhigh_tier(opus_4_7(), "xhigh", "xhigh")] - #[case::max(opus_4_7(), "max", "max")] - #[case::minimal_maps_to_low_on_4_6(sonnet_4_6(), "minimal", "low")] - #[case::low_on_4_6(sonnet_4_6(), "low", "low")] - #[case::max_without_max_tier_is_allowed_on_adaptive_models(newfamily_6(), "max", "max")] - fn reasoning_effort_on_adaptive_model_becomes_summarized_adaptive_thinking_and_effort( - #[case] capabilities: AnthropicModelCapabilities, - #[case] reasoning_effort: &str, - #[case] expected_effort: &str, - ) { - assert_eq!( - translate( - capabilities, - json!({"max_tokens": 1024, "reasoning_effort": reasoning_effort}) - ), - Ok(request(json!({ - "max_tokens": 1024, - "thinking": {"type": "adaptive", "display": "summarized"}, - "output_config": {"effort": expected_effort} - }))) - ); - } - - #[rstest] - #[case::adaptive_shape_is_not_dropped_for_small_max_tokens( - opus_4_7(), - json!({"max_tokens": 64, "reasoning_effort": "high"}), - json!({"max_tokens": 64, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}}) - )] - #[case::caller_output_config_effort_wins( - opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": "low", "output_config": {"effort": "max"}}), - json!({"max_tokens": 1024, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "max"}}) - )] - #[case::effort_merges_into_caller_output_config( - opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": "high", "output_config": {"format": {"type": "json_schema"}}}), - json!({ - "max_tokens": 1024, - "thinking": {"type": "adaptive", "display": "summarized"}, - "output_config": {"format": {"type": "json_schema"}, "effort": "high"} - }) - )] - #[case::non_object_output_config_is_replaced( - opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": "high", "output_config": "bogus"}), - json!({"max_tokens": 1024, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}}) - )] - #[case::caller_thinking_and_output_config_win( - sonnet_4_6(), - json!({ - "max_tokens": 16000, - "reasoning_effort": "low", - "thinking": {"type": "enabled", "budget_tokens": 8000}, - "output_config": {"effort": "high"} - }), - json!({ - "max_tokens": 16000, - "thinking": {"type": "enabled", "budget_tokens": 8000}, - "output_config": {"effort": "high"} - }) - )] - #[case::caller_legacy_thinking_is_then_translated_while_reasoning_effort_level_stays( - opus_4_7(), - json!({"max_tokens": 16000, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), - json!({"max_tokens": 16000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "low"}}) - )] - #[case::caller_disabled_thinking_is_kept_then_omitted_on_always_on_model( - fable_5_1(), - json!({"max_tokens": 1024, "reasoning_effort": "high", "thinking": {"type": "disabled"}}), - json!({"max_tokens": 1024, "output_config": {"effort": "high"}}) - )] - #[case::non_adaptive_model_gets_no_output_config( - opus_4_5(), - json!({"max_tokens": 8192, "reasoning_effort": "high"}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) - )] - #[case::caller_thinking_wins_on_non_adaptive_model( - opus_4_5(), - json!({"max_tokens": 16000, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), - json!({"max_tokens": 16000, "thinking": {"type": "enabled", "budget_tokens": 8000}}) - )] - #[case::caller_thinking_survives_when_mapped_budget_cannot_fit( - opus_4_5(), - json!({"max_tokens": 1024, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 8000}}) - )] - #[case::missing_max_tokens_leaves_budget_uncapped( - haiku_4_5(), - json!({"reasoning_effort": "high"}), - json!({"thinking": {"type": "enabled", "budget_tokens": 4096}}) - )] - #[case::budget_below_max_tokens_is_kept( - haiku_4_5(), - json!({"max_tokens": 4097, "reasoning_effort": "high"}), - json!({"max_tokens": 4097, "thinking": {"type": "enabled", "budget_tokens": 4096}}) - )] - #[case::budget_equal_to_max_tokens_is_capped( - haiku_4_5(), - json!({"max_tokens": 4096, "reasoning_effort": "high"}), - json!({"max_tokens": 4096, "thinking": {"type": "enabled", "budget_tokens": 4095}}) - )] - #[case::budget_above_max_tokens_is_capped( - haiku_4_5(), - json!({"max_tokens": 4000, "reasoning_effort": "xhigh"}), - json!({"max_tokens": 4000, "thinking": {"type": "enabled", "budget_tokens": 3999}}) - )] - #[case::max_tokens_just_above_min_budget_caps_to_min_budget( - haiku_4_5(), - json!({"max_tokens": 1025, "reasoning_effort": "xhigh"}), - json!({"max_tokens": 1025, "thinking": {"type": "enabled", "budget_tokens": 1024}}) - )] - #[case::max_tokens_at_min_budget_drops_thinking( - haiku_4_5(), - json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), - json!({"max_tokens": 1024}) - )] - #[case::pinned_temperature_is_dropped_after_thinking_is_synthesized( - haiku_4_5(), - json!({"max_tokens": 8192, "reasoning_effort": "low", "temperature": 0}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 1024}}) - )] - fn reasoning_effort_is_translated( - #[case] capabilities: AnthropicModelCapabilities, - #[case] input: Value, - #[case] expected: Value, - ) { - assert_eq!(translate(capabilities, input), Ok(request(expected))); - } - - #[rstest] - #[case::minimal_floors_at_min_budget("minimal", 1024)] - #[case::low("low", 1024)] - #[case::medium("medium", 2048)] - #[case::high("high", 4096)] - #[case::xhigh("xhigh", 8192)] - #[case::max("max", 16384)] - fn reasoning_effort_on_non_adaptive_model_uses_the_tier_budget( - haiku_4_5: AnthropicModelCapabilities, - #[case] reasoning_effort: &str, - #[case] expected_budget: u64, - ) { - assert_eq!( - translate( - haiku_4_5, - json!({"max_tokens": 32000, "reasoning_effort": reasoning_effort}) - ), - Ok(request(json!({ - "max_tokens": 32000, - "thinking": {"type": "enabled", "budget_tokens": expected_budget} - }))) - ); - } - - #[rstest] - #[case::adaptive_model(opus_4_7())] - #[case::effort_capable_model(opus_4_5())] - #[case::budget_model(haiku_4_5())] - fn reasoning_effort_none_clears_thinking_and_output_config( - #[case] capabilities: AnthropicModelCapabilities, - ) { - assert_eq!( - translate( - capabilities, - json!({ - "max_tokens": 1024, - "reasoning_effort": "none", - "thinking": {"type": "adaptive"}, - "output_config": {"effort": "high"} - }) - ), - Ok(request(json!({"max_tokens": 1024}))) - ); - } - - #[rstest] - #[case::bogus_on_budget_model( - opus_4_5(), - json!({"max_tokens": 1024, "reasoning_effort": "bogus"}), - format!("Unmapped reasoning effort: 'bogus'. Must be one of: {EFFORT_CHOICES}.") - )] - #[case::disabled_on_budget_model( - haiku_4_5(), - json!({"max_tokens": 1024, "reasoning_effort": "disabled"}), - format!("Unmapped reasoning effort: 'disabled'. Must be one of: {EFFORT_CHOICES}.") - )] - #[case::empty_on_budget_model( - haiku_4_5(), - json!({"max_tokens": 1024, "reasoning_effort": ""}), - format!("Unmapped reasoning effort: ''. Must be one of: {EFFORT_CHOICES}.") - )] - #[case::invalid_on_adaptive_model( - opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": "invalid"}), - format!("Invalid reasoning_effort: 'invalid'. Must be one of: {EFFORT_CHOICES}") - )] - #[case::disabled_on_adaptive_model( - opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": "disabled"}), - format!("Invalid reasoning_effort: 'disabled'. Must be one of: {EFFORT_CHOICES}") - )] - #[case::empty_on_adaptive_model( - opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": ""}), - format!("Invalid reasoning_effort: ''. Must be one of: {EFFORT_CHOICES}") - )] - #[case::xhigh_without_xhigh_tier_on_4_6( - sonnet_4_6(), - json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), - "effort='xhigh' is not supported by this model. Got model: claude".to_string() - )] - #[case::xhigh_without_xhigh_tier_on_unmapped_adaptive_model( - newfamily_6(), - json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), - "effort='xhigh' is not supported by this model. Got model: claude".to_string() - )] - #[case::unrecognized_adaptive_effort_on_budget_model( - haiku_4_5(), - claude_code_payload("turbo", 8192), - format!("Unmapped reasoning effort: 'turbo'. Must be one of: {EFFORT_CHOICES}.") - )] - fn unsupported_effort_is_a_request_error( - #[case] capabilities: AnthropicModelCapabilities, - #[case] input: Value, - #[case] expected_message: String, - ) { - assert_eq!( - translate(capabilities, input), - Err(Error::InvalidRequest(expected_message)) - ); - } - - #[rstest] - #[case::omitted_on_always_on_model(fable_5_1(), json!({"type": "disabled"}), None)] - #[case::kept_on_adaptive_model(opus_4_7(), json!({"type": "disabled"}), Some(json!({"type": "disabled"})))] - #[case::kept_on_budget_model(haiku_4_5(), json!({"type": "disabled"}), Some(json!({"type": "disabled"})))] - #[case::adaptive_kept_on_always_on_model( - fable_5_1(), - json!({"type": "adaptive"}), - Some(json!({"type": "adaptive"})) - )] - fn disabled_thinking_is_omitted_only_for_always_on_models( - #[case] capabilities: AnthropicModelCapabilities, - #[case] thinking: Value, - #[case] expected_thinking: Option, - ) { - let expected = match expected_thinking { - Some(thinking) => json!({"max_tokens": 64, "thinking": thinking}), - None => json!({"max_tokens": 64}), - }; - assert_eq!( - translate( - capabilities, - json!({"max_tokens": 64, "thinking": thinking}) - ), - Ok(request(expected)) - ); - } - - #[rstest] - #[case::far_above_xhigh_budget(opus_4_7(), json!(16384), "xhigh")] - #[case::at_xhigh_budget(opus_4_7(), json!(8192), "xhigh")] - #[case::below_xhigh_budget(opus_4_7(), json!(8191), "high")] - #[case::xhigh_budget_without_xhigh_tier(newfamily_6(), json!(8192), "high")] - #[case::large_budget_without_xhigh_tier(newfamily_6(), json!(31999), "high")] - #[case::at_high_budget(opus_4_7(), json!(4096), "high")] - #[case::below_high_budget(opus_4_7(), json!(4095), "medium")] - #[case::at_medium_budget(opus_4_7(), json!(2048), "medium")] - #[case::below_medium_budget(opus_4_7(), json!(2047), "low")] - #[case::tiny_budget(opus_4_7(), json!(1), "low")] - #[case::missing_budget(opus_4_7(), Value::Null, "low")] - #[case::always_on_model(fable_5_1(), json!(24000), "xhigh")] - fn legacy_thinking_is_bucketed_into_adaptive_effort_on_adaptive_only_models( - #[case] capabilities: AnthropicModelCapabilities, - #[case] budget_tokens: Value, - #[case] expected_effort: &str, - ) { - let thinking = match budget_tokens { - Value::Null => json!({"type": "enabled"}), - budget_tokens => json!({"type": "enabled", "budget_tokens": budget_tokens}), - }; - assert_eq!( - translate( - capabilities, - json!({"max_tokens": 1024, "thinking": thinking}) - ), - Ok(request(json!({ - "max_tokens": 1024, - "thinking": {"type": "adaptive"}, - "output_config": {"effort": expected_effort} - }))) - ); - } - - #[rstest] - #[case::verbatim_on_model_accepting_legacy_thinking( - sonnet_4_6(), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}) - )] - #[case::verbatim_with_explicit_output_config_on_model_accepting_legacy_thinking( - sonnet_4_6(), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}, "output_config": {"effort": "low"}}), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}, "output_config": {"effort": "low"}}) - )] - #[case::verbatim_on_non_adaptive_model( - opus_4_5(), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}) - )] - #[case::caller_output_config_effort_wins( - opus_4_7(), - json!({ - "max_tokens": 32000, - "thinking": {"type": "enabled", "budget_tokens": 31999}, - "output_config": {"effort": "low", "format": {"type": "json_schema"}} - }), - json!({ - "max_tokens": 32000, - "thinking": {"type": "adaptive"}, - "output_config": {"effort": "low", "format": {"type": "json_schema"}} - }) - )] - #[case::effort_merges_into_caller_output_config( - opus_4_7(), - json!({ - "max_tokens": 32000, - "thinking": {"type": "enabled", "budget_tokens": 4096}, - "output_config": {"format": {"type": "json_schema"}} - }), - json!({ - "max_tokens": 32000, - "thinking": {"type": "adaptive"}, - "output_config": {"effort": "high", "format": {"type": "json_schema"}} - }) - )] - #[case::adaptive_thinking_is_left_alone( - opus_4_7(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive", "display": "summarized"}}), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive", "display": "summarized"}}) - )] - fn legacy_thinking_on_adaptive_capable_models( - #[case] capabilities: AnthropicModelCapabilities, - #[case] input: Value, - #[case] expected: Value, - ) { - assert_eq!(translate(capabilities, input), Ok(request(expected))); - } - - #[rstest] - #[case::bare_adaptive_becomes_medium_budget_on_budget_model( - haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) - )] - #[case::medium_effort_becomes_medium_budget_on_budget_model( - haiku_4_5(), - claude_code_payload("medium", 8192), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) - )] - #[case::empty_effort_becomes_medium_budget_on_budget_model( - haiku_4_5(), - claude_code_payload("", 8192), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) - )] - #[case::high_effort_becomes_high_budget_on_budget_model( - haiku_4_5(), - claude_code_payload("high", 8192), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) - )] - #[case::effort_only_becomes_budget_on_budget_model( - haiku_4_5(), - json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) - )] - #[case::effort_replaces_caller_legacy_budget_on_budget_model( - haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 3000}, "output_config": {"effort": "high"}}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) - )] - #[case::residual_output_config_survives_effort_translation( - haiku_4_5(), - json!({ - "max_tokens": 8192, - "thinking": {"type": "adaptive"}, - "output_config": {"effort": "medium", "format": {"type": "json_schema"}} - }), - json!({ - "max_tokens": 8192, - "thinking": {"type": "enabled", "budget_tokens": 2048}, - "output_config": {"format": {"type": "json_schema"}} - }) - )] - #[case::effortless_output_config_is_kept( - haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {"format": {"type": "json_schema"}}}), - json!({ - "max_tokens": 8192, - "thinking": {"type": "enabled", "budget_tokens": 2048}, - "output_config": {"format": {"type": "json_schema"}} - }) - )] - #[case::empty_output_config_is_kept( - haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {}}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}, "output_config": {}}) - )] - #[case::missing_max_tokens_leaves_budget_uncapped( - haiku_4_5(), - json!({"thinking": {"type": "adaptive"}}), - json!({"thinking": {"type": "enabled", "budget_tokens": 2048}}) - )] - #[case::budget_is_capped_below_max_tokens( - haiku_4_5(), - claude_code_payload("high", 3000), - json!({"max_tokens": 3000, "thinking": {"type": "enabled", "budget_tokens": 2999}}) - )] - #[case::max_tokens_just_above_min_budget_caps_to_min_budget( - haiku_4_5(), - claude_code_payload("medium", 1025), - json!({"max_tokens": 1025, "thinking": {"type": "enabled", "budget_tokens": 1024}}) - )] - #[case::max_tokens_at_min_budget_drops_thinking_and_effort( - haiku_4_5(), - claude_code_payload("medium", 1024), - json!({"max_tokens": 1024}) - )] - #[case::max_tokens_below_min_budget_drops_thinking_and_effort( - haiku_4_5(), - claude_code_payload("medium", 512), - json!({"max_tokens": 512}) - )] - #[case::bare_adaptive_is_dropped_on_non_reasoning_model( - haiku_3_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), - json!({"max_tokens": 8192}) - )] - #[case::adaptive_and_effort_are_dropped_on_non_reasoning_model( - haiku_3_5(), - claude_code_payload("medium", 8192), - json!({"max_tokens": 8192}) - )] - #[case::effort_only_is_dropped_on_non_reasoning_model( - haiku_3_5(), - json!({"max_tokens": 8192, "output_config": {"effort": "high", "format": {"type": "json_schema"}}}), - json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}) - )] - #[case::supported_effort_is_kept_and_adaptive_thinking_dropped_on_effort_model( - opus_4_5(), - claude_code_payload("medium", 8192), - json!({"max_tokens": 8192, "output_config": {"effort": "medium"}}) - )] - #[case::bare_adaptive_is_dropped_on_effort_model( - opus_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), - json!({"max_tokens": 8192}) - )] - #[case::effort_only_is_left_alone_on_effort_model( - opus_4_5(), - json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), - json!({"max_tokens": 8192, "output_config": {"effort": "high"}}) - )] - #[case::unsupported_effort_only_is_left_for_provider_normalization( - opus_4_5(), - json!({"max_tokens": 4096, "output_config": {"effort": "xhigh"}}), - json!({"max_tokens": 4096, "output_config": {"effort": "xhigh"}}) - )] - #[case::legacy_thinking_is_kept_beside_native_effort_on_effort_model( - opus_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}, "output_config": {"effort": "high"}}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}, "output_config": {"effort": "high"}}) - )] - #[case::unsupported_xhigh_with_adaptive_thinking_falls_back_to_budget( - opus_4_5(), - claude_code_payload("xhigh", 64000), - json!({"max_tokens": 64000, "thinking": {"type": "enabled", "budget_tokens": 8192}}) - )] - #[case::unsupported_max_with_adaptive_thinking_falls_back_to_budget( - opus_4_5(), - claude_code_payload("max", 64000), - json!({"max_tokens": 64000, "thinking": {"type": "enabled", "budget_tokens": 16384}}) - )] - #[case::bare_adaptive_is_native_on_4_6( - sonnet_4_6(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}) - )] - #[case::adaptive_payload_is_native_on_4_6( - sonnet_4_6(), - claude_code_payload("high", 8192), - claude_code_payload("high", 8192) - )] - #[case::request_without_adaptive_interface_is_left_alone( - haiku_4_5(), - json!({"max_tokens": 1024}), - json!({"max_tokens": 1024}) - )] - fn adaptive_interface_is_reshaped_for_non_adaptive_models( - #[case] capabilities: AnthropicModelCapabilities, - #[case] input: Value, - #[case] expected: Value, - ) { - assert_eq!(translate(capabilities, input), Ok(request(expected))); - } - - #[rstest] - #[case::adaptive_downgraded_to_enabled_thinking( - haiku_4_5(), - claude_code_payload("medium", 8192), - 0.0, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) - )] - #[case::bare_adaptive_downgraded_to_enabled_thinking( - haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), - 0.0, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) - )] - #[case::reasoning_effort_synthesized_enabled_thinking( - haiku_4_5(), - json!({"max_tokens": 8192, "reasoning_effort": "high"}), - 0.2, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) - )] - #[case::above_one_with_enabled_thinking( - haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}), - 1.5, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) - )] - #[case::native_effort_kept_on_effort_model( - opus_4_5(), - claude_code_payload("medium", 8192), - 0.0, - json!({"max_tokens": 8192, "output_config": {"effort": "medium"}}) - )] - #[case::effort_only_on_effort_model( - opus_4_5(), - json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), - 0.0, - json!({"max_tokens": 8192, "output_config": {"effort": "high"}}) - )] - fn pinned_temperature_is_dropped_when_thinking_or_effort_survives_on_non_adaptive_model( - #[case] capabilities: AnthropicModelCapabilities, - #[case] input: Value, - #[case] temperature: f64, - #[case] expected: Value, - ) { - assert_eq!( - translate(capabilities, with_temperature(input, temperature)), - Ok(request(expected)) - ); - } - - #[rstest] - #[case::temperature_one_with_enabled_thinking( - haiku_4_5(), - claude_code_payload("medium", 8192), - 1.0, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) - )] - #[case::thinking_dropped_for_small_max_tokens( - haiku_4_5(), - claude_code_payload("medium", 512), - 0.0, - json!({"max_tokens": 512}) - )] - #[case::thinking_dropped_on_non_reasoning_model( - haiku_3_5(), - claude_code_payload("medium", 8192), - 0.0, - json!({"max_tokens": 8192}) - )] - #[case::disabled_thinking( - haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "disabled"}}), - 0.0, - json!({"max_tokens": 8192, "thinking": {"type": "disabled"}}) - )] - #[case::no_thinking(haiku_4_5(), json!({"max_tokens": 8192}), 0.0, json!({"max_tokens": 8192}))] - #[case::output_config_without_effort( - haiku_4_5(), - json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}), - 0.0, - json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}) - )] - #[case::adaptive_model( - opus_4_7(), - claude_code_payload("medium", 8192), - 0.0, - claude_code_payload("medium", 8192) - )] - #[case::legacy_thinking_on_adaptive_model( - sonnet_4_6(), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}), - 0.0, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) - )] - fn temperature_is_kept( - #[case] capabilities: AnthropicModelCapabilities, - #[case] input: Value, - #[case] temperature: f64, - #[case] expected: Value, - ) { - assert_eq!( - translate(capabilities, with_temperature(input, temperature)), - Ok(request(with_temperature(expected, temperature))) - ); - } - - #[rstest] - #[case::minimal("MINIMAL", ThinkingBudgets { minimal: 5000, ..ThinkingBudgets::default() })] - #[case::low("LOW", ThinkingBudgets { low: 5000, ..ThinkingBudgets::default() })] - #[case::medium("MEDIUM", ThinkingBudgets { medium: 5000, ..ThinkingBudgets::default() })] - #[case::high("HIGH", ThinkingBudgets { high: 5000, ..ThinkingBudgets::default() })] - #[case::xhigh("XHIGH", ThinkingBudgets { xhigh: 5000, ..ThinkingBudgets::default() })] - #[case::max("MAX", ThinkingBudgets { max: 5000, ..ThinkingBudgets::default() })] - fn each_tier_budget_reads_only_its_own_environment_override( - #[case] tier: &str, - #[case] expected: ThinkingBudgets, - ) { - assert_eq!(overridden_budgets(&[(tier, "5000")]), expected); - } - - #[rstest] - #[case::whitespace_is_trimmed(" 6000 ", 6000)] - #[case::unparseable_value_keeps_default("lots", 4096)] - fn environment_override_parsing(#[case] raw: &str, #[case] expected_high: u64) { - assert_eq!( - overridden_budgets(&[("HIGH", raw)]), - ThinkingBudgets { - high: expected_high, - ..ThinkingBudgets::default() - } - ); - } - - #[rstest] - #[case::reasoning_effort_uses_overridden_budget( - &[("HIGH", "6000")], - haiku_4_5(), - json!({"max_tokens": 32000, "reasoning_effort": "high"}), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 6000}}) - )] - #[case::minimal_override_below_min_budget_is_floored( - &[("MINIMAL", "512")], - haiku_4_5(), - json!({"max_tokens": 32000, "reasoning_effort": "minimal"}), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 1024}}) - )] - #[case::minimal_override_above_min_budget_is_used( - &[("MINIMAL", "2000")], - haiku_4_5(), - json!({"max_tokens": 32000, "reasoning_effort": "minimal"}), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 2000}}) - )] - #[case::adaptive_fallback_uses_overridden_medium_budget( - &[("MEDIUM", "3000")], - haiku_4_5(), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}}), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 3000}}) - )] - #[case::legacy_bucket_below_overridden_high_budget( - &[("HIGH", "6000")], - opus_4_7(), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 5999}}), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}}) - )] - #[case::legacy_bucket_at_overridden_high_budget( - &[("HIGH", "6000")], - opus_4_7(), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 6000}}), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}) - )] - #[case::legacy_bucket_below_overridden_xhigh_budget( - &[("XHIGH", "20000")], - opus_4_7(), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 19999}}), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}) - )] - #[case::legacy_bucket_at_overridden_medium_budget( - &[("MEDIUM", "3000")], - opus_4_7(), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 3000}}), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}}) - )] - #[case::legacy_bucket_below_overridden_medium_budget( - &[("MEDIUM", "3000")], - opus_4_7(), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 2999}}), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "low"}}) - )] - fn translation_honors_budget_overrides( - #[case] overrides: &[(&str, &str)], - #[case] capabilities: AnthropicModelCapabilities, - #[case] input: Value, - #[case] expected: Value, - ) { - let context = ThinkingContext { - capabilities, - budgets: overridden_budgets(overrides), - }; - assert_eq!( - translate_thinking(request(input), &context), - Ok(request(expected)) - ); - } -} diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/mod.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/mod.rs deleted file mode 100644 index ba63992f3cb..00000000000 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod messages; diff --git a/litellm-rust/crates/llms/src/anthropic/messages/AGENTS.md b/litellm-rust/crates/llms/src/anthropic/messages/AGENTS.md new file mode 100644 index 00000000000..b7832c4b8f3 --- /dev/null +++ b/litellm-rust/crates/llms/src/anthropic/messages/AGENTS.md @@ -0,0 +1 @@ +- https://platform.claude.com/docs/en/api/http/messages/create diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/handler.rs b/litellm-rust/crates/llms/src/anthropic/messages/handler.rs similarity index 82% rename from litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/handler.rs rename to litellm-rust/crates/llms/src/anthropic/messages/handler.rs index 0e2ab97956a..9e187035944 100644 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/handler.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/handler.rs @@ -1,14 +1,18 @@ -use litellm_types::llms::anthropic_messages::anthropic_request::{ - AnthropicMessage, AnthropicMessagesRequest, +use litellm_types::{ + llms::anthropic_messages::anthropic_request::{ + AdaptiveThinking, AnthropicMessage, AnthropicMessagesOptionalParams, + AnthropicMessagesRequest, EnabledThinking, ThinkingConfig, ThinkingDisplay, + }, + recognized::Recognized, }; use serde_json::{Value, json}; use crate::{ + Error, anthropic::common_utils::{ flatten_unencrypted_web_search_results, sanitize_tool_use_ids, strip_empty_content_blocks, strip_provider_specific_fields, }, - base_llm::chat::transformation::Error, }; pub fn shape_anthropic_messages_request( @@ -17,12 +21,16 @@ pub fn shape_anthropic_messages_request( ) -> Result { Ok(AnthropicMessagesRequest { messages: sanitize_anthropic_messages(request.messages), - metadata: request - .metadata - .as_ref() - .map(validate_anthropic_api_metadata) - .transpose()?, - thinking: with_reasoning_auto_summary(request.thinking, reasoning_auto_summary), + params: AnthropicMessagesOptionalParams { + metadata: request + .params + .metadata + .as_ref() + .map(validate_anthropic_api_metadata) + .transpose()?, + thinking: with_reasoning_auto_summary(request.params.thinking, reasoning_auto_summary), + ..request.params + }, ..request }) } @@ -48,20 +56,38 @@ fn validate_anthropic_api_metadata(metadata: &Value) -> Result { } } -fn with_reasoning_auto_summary(thinking: Option, enabled: bool) -> Option { - let Some(Value::Object(thinking)) = thinking else { +fn with_reasoning_auto_summary( + thinking: Option>, + enabled: bool, +) -> Option> { + if !enabled { return thinking; - }; - if !enabled || thinking.get("type").and_then(Value::as_str) == Some("disabled") { - return Some(Value::Object(thinking)); } - Some(Value::Object( - thinking - .into_iter() - .filter(|(key, _)| key != "display") - .chain([("display".to_string(), json!("summarized"))]) - .collect(), - )) + let summarized = Some(Recognized::Known(ThinkingDisplay::Summarized)); + match thinking { + Some(Recognized::Known(ThinkingConfig::Enabled(enabled))) => Some(Recognized::Known( + ThinkingConfig::Enabled(EnabledThinking { + display: summarized, + ..enabled + }), + )), + Some(Recognized::Known(ThinkingConfig::Adaptive(adaptive))) => Some(Recognized::Known( + ThinkingConfig::Adaptive(AdaptiveThinking { + display: summarized, + ..adaptive + }), + )), + Some(Recognized::Unrecognized(Value::Object(fields))) => { + Some(Recognized::Unrecognized(Value::Object( + fields + .into_iter() + .filter(|(key, _)| key != "display") + .chain([("display".to_string(), json!("summarized"))]) + .collect(), + ))) + } + other => other, + } } #[cfg(test)] @@ -230,12 +256,22 @@ mod tests { )] #[case::no_thinking(None, true, None)] #[case::non_object_thinking(Some(json!("enabled")), true, Some(json!("enabled")))] + #[case::unknown_type( + Some(json!({"type": "future"})), + true, + Some(json!({"type": "future", "display": "summarized"})), + )] fn reasoning_auto_summary_marks_active_thinking_as_summarized( #[case] thinking: Option, #[case] enabled: bool, #[case] expected: Option, ) { - assert_eq!(with_reasoning_auto_summary(thinking, enabled), expected); + let thinking = thinking.map(|thinking| serde_json::from_value(thinking).unwrap()); + assert_eq!( + with_reasoning_auto_summary(thinking, enabled) + .map(|thinking| serde_json::to_value(thinking).unwrap()), + expected + ); } #[test] diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/mod.rs b/litellm-rust/crates/llms/src/anthropic/messages/mod.rs similarity index 83% rename from litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/mod.rs rename to litellm-rust/crates/llms/src/anthropic/messages/mod.rs index 5adf5fda16f..dff4bf18bd5 100644 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/mod.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/mod.rs @@ -1,5 +1,4 @@ pub mod handler; -pub mod headers; pub mod streaming_iterator; pub mod thinking; pub mod transformation; diff --git a/litellm-rust/crates/llms/src/anthropic/messages/streaming_iterator.rs b/litellm-rust/crates/llms/src/anthropic/messages/streaming_iterator.rs new file mode 100644 index 00000000000..4f00cd0af7e --- /dev/null +++ b/litellm-rust/crates/llms/src/anthropic/messages/streaming_iterator.rs @@ -0,0 +1,142 @@ +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct AnthropicStreamUsage { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub input_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub output_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cache_creation_input_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cache_read_input_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub server_tool_use: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct AnthropicStreamMessage { + pub id: String, + #[serde(rename = "type")] + pub message_type: String, + pub role: String, + pub model: String, + pub content: Vec, + pub stop_reason: Option, + pub stop_sequence: Option, + pub usage: AnthropicStreamUsage, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum AnthropicContentBlockDelta { + TextDelta { + text: String, + }, + InputJsonDelta { + partial_json: String, + }, + #[serde(rename = "citations_delta")] + Citations { + citation: Value, + }, + ThinkingDelta { + thinking: String, + }, + SignatureDelta { + signature: String, + }, + CompactionDelta { + content: String, + }, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct AnthropicContentBlock { + #[serde(rename = "type")] + pub block_type: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub text: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub input: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub thinking: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub signature: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub data: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub caller: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct AnthropicMessageDelta { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub stop_reason: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub stop_sequence: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub stop_details: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub container: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct AnthropicStreamError { + #[serde(rename = "type")] + pub error_type: String, + pub message: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub details: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum AnthropicMessagesStreamEvent { + MessageStart { + message: AnthropicStreamMessage, + }, + ContentBlockStart { + index: u64, + content_block: AnthropicContentBlock, + }, + ContentBlockDelta { + index: u64, + delta: AnthropicContentBlockDelta, + }, + ContentBlockStop { + index: u64, + }, + MessageDelta { + delta: AnthropicMessageDelta, + #[serde(default, skip_serializing_if = "Option::is_none")] + usage: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + context_management: Option, + }, + MessageStop { + #[serde(default, skip_serializing_if = "Option::is_none")] + usage: Option, + }, + Ping, + Error { + error: AnthropicStreamError, + }, +} diff --git a/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs new file mode 100644 index 00000000000..101ed438738 --- /dev/null +++ b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs @@ -0,0 +1,1292 @@ +use litellm_core_utils::settings::Lookup; +use litellm_python_compat::{json::from_json, repr::repr, truthy::truthy}; +use litellm_types::{ + llms::{ + anthropic_messages::anthropic_request::{ + AnthropicMessagesOptionalParams, AnthropicMessagesRequest, EffortLevel, OutputConfig, + ThinkingConfig, ThinkingDisplay, + }, + openai::ReasoningEffort, + }, + recognized::Recognized, +}; +use serde_json::Value; + +use crate::{Error, anthropic::common_utils::AnthropicModelCapabilities}; + +pub const ANTHROPIC_MIN_THINKING_BUDGET_TOKENS: u64 = 1024; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ThinkingBudgets { + pub minimal: u64, + pub low: u64, + pub medium: u64, + pub high: u64, + pub xhigh: u64, + pub max: u64, +} + +impl Default for ThinkingBudgets { + fn default() -> Self { + Self { + minimal: 128, + low: 1024, + medium: 2048, + high: 4096, + xhigh: 8192, + max: 16384, + } + } +} + +impl ThinkingBudgets { + pub fn from_lookup(env: &impl Lookup) -> Self { + let defaults = Self::default(); + let tier = |name: &str, default: u64| { + env.parsed::(&format!("DEFAULT_REASONING_EFFORT_{name}_THINKING_BUDGET")) + .unwrap_or(default) + }; + Self { + minimal: tier("MINIMAL", defaults.minimal), + low: tier("LOW", defaults.low), + medium: tier("MEDIUM", defaults.medium), + high: tier("HIGH", defaults.high), + xhigh: tier("XHIGH", defaults.xhigh), + max: tier("MAX", defaults.max), + } + } + + fn for_effort(&self, effort: ReasoningEffort) -> Option { + match effort { + ReasoningEffort::None => None, + ReasoningEffort::Minimal => { + Some(self.minimal.max(ANTHROPIC_MIN_THINKING_BUDGET_TOKENS)) + } + ReasoningEffort::Low => Some(self.low), + ReasoningEffort::Medium => Some(self.medium), + ReasoningEffort::High => Some(self.high), + ReasoningEffort::Xhigh => Some(self.xhigh), + ReasoningEffort::Max => Some(self.max), + } + } + + fn effort_for_budget( + &self, + budget_tokens: u64, + capabilities: &AnthropicModelCapabilities, + ) -> EffortLevel { + if budget_tokens >= self.xhigh && capabilities.effort_tiers.xhigh { + return EffortLevel::Xhigh; + } + if budget_tokens >= self.high { + return EffortLevel::High; + } + if budget_tokens >= self.medium { + return EffortLevel::Medium; + } + EffortLevel::Low + } +} + +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub struct ThinkingContext { + pub capabilities: AnthropicModelCapabilities, + pub budgets: ThinkingBudgets, +} + +fn unmapped_effort(effort: &Value) -> Error { + let choices = ReasoningEffort::ALL + .map(|effort| format!("'{}'", effort.as_str())) + .join(", "); + Error::InvalidRequest(format!( + "Unmapped reasoning effort: {}. Must be one of: {choices}.", + repr(&from_json(effort.clone())) + )) +} + +fn unsupported_effort(level: EffortLevel, model: &str) -> Error { + Error::InvalidRequest(format!( + "effort='{}' is not supported by this model. Got model: {model}", + level.as_str() + )) +} + +fn output_effort(effort: ReasoningEffort) -> Option { + match effort { + ReasoningEffort::None => None, + ReasoningEffort::Minimal | ReasoningEffort::Low => Some(EffortLevel::Low), + ReasoningEffort::Medium => Some(EffortLevel::Medium), + ReasoningEffort::High => Some(EffortLevel::High), + ReasoningEffort::Xhigh => Some(EffortLevel::Xhigh), + ReasoningEffort::Max => Some(EffortLevel::Max), + } +} + +fn fit_budget_to_max_tokens(budget_tokens: u64, max_tokens: Option) -> Option { + let Some(max_tokens) = max_tokens else { + return Some(budget_tokens); + }; + (max_tokens > ANTHROPIC_MIN_THINKING_BUDGET_TOKENS).then(|| budget_tokens.min(max_tokens - 1)) +} + +fn known_thinking(request: &AnthropicMessagesRequest) -> Option<&ThinkingConfig> { + request.params.thinking.as_ref().and_then(Recognized::known) +} + +fn known_effort(request: &AnthropicMessagesRequest) -> Option<&Recognized> { + request + .params + .output_config + .as_ref() + .and_then(Recognized::known) + .and_then(|config| config.effort.as_ref()) +} + +fn with_default_effort( + output_config: Option>, + level: EffortLevel, +) -> Option> { + let config = match output_config { + Some(Recognized::Known(config)) => config, + _ => OutputConfig::default(), + }; + Some(Recognized::Known(OutputConfig { + effort: Some(config.effort.unwrap_or(Recognized::Known(level))), + ..config + })) +} + +fn without_effort( + output_config: Option>, +) -> Option> { + let Some(Recognized::Known(config)) = output_config else { + return output_config; + }; + if config.effort.is_none() { + return Some(Recognized::Known(config)); + } + let residual = OutputConfig { + effort: None, + ..config + }; + (!residual.is_empty()).then_some(Recognized::Known(residual)) +} + +fn legacy_reasoning_effort( + effort: Option<&Recognized>, +) -> Result { + match effort { + Some(Recognized::Known(level)) => Ok((*level).into()), + Some(Recognized::Unrecognized(value)) if truthy(&from_json(value.clone())) => value + .as_str() + .and_then(ReasoningEffort::parse) + .ok_or_else(|| unmapped_effort(value)), + None | Some(Recognized::Unrecognized(_)) => Ok(ReasoningEffort::Medium), + } +} + +fn translate_reasoning_effort( + request: AnthropicMessagesRequest, + context: &ThinkingContext, +) -> Result { + let Some(reasoning_effort) = request.params.reasoning_effort else { + return Ok(request); + }; + let request = AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { + reasoning_effort: None, + ..request.params + }, + ..request + }; + let effort = match reasoning_effort { + Recognized::Known(effort) => effort, + Recognized::Unrecognized(value @ Value::String(_)) => { + return Err(unmapped_effort(&value)); + } + Recognized::Unrecognized(_) => return Ok(request), + }; + let (Some(level), Some(budget)) = (output_effort(effort), context.budgets.for_effort(effort)) + else { + return Ok(AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { + thinking: None, + output_config: None, + ..request.params + }, + ..request + }); + }; + let capabilities = &context.capabilities; + if capabilities.supports_adaptive_thinking { + if !capabilities.accepts_effort(level) { + return Err(unsupported_effort(level, &request.model)); + } + let adaptive = ThinkingConfig::adaptive(Some(ThinkingDisplay::Summarized)); + return Ok(AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { + thinking: Some( + request + .params + .thinking + .unwrap_or(Recognized::Known(adaptive)), + ), + output_config: with_default_effort(request.params.output_config, level), + ..request.params + }, + ..request + }); + } + let Some(budget) = fit_budget_to_max_tokens(budget, request.params.max_tokens) else { + return Ok(request); + }; + let enabled = ThinkingConfig::enabled(budget); + Ok(AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { + thinking: Some( + request + .params + .thinking + .unwrap_or(Recognized::Known(enabled)), + ), + ..request.params + }, + ..request + }) +} + +fn drop_disabled_thinking( + request: AnthropicMessagesRequest, + context: &ThinkingContext, +) -> AnthropicMessagesRequest { + if !context.capabilities.thinking_always_on + || !matches!(known_thinking(&request), Some(ThinkingConfig::Disabled(_))) + { + return request; + } + AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { + thinking: None, + ..request.params + }, + ..request + } +} + +fn translate_legacy_thinking_for_adaptive_model( + request: AnthropicMessagesRequest, + context: &ThinkingContext, +) -> AnthropicMessagesRequest { + let capabilities = &context.capabilities; + if !capabilities.supports_adaptive_thinking || capabilities.supports_legacy_thinking { + return request; + } + let Some(ThinkingConfig::Enabled(enabled)) = known_thinking(&request) else { + return request; + }; + let budget = enabled + .budget_tokens + .as_ref() + .and_then(Recognized::known) + .copied() + .unwrap_or(0); + let level = context.budgets.effort_for_budget(budget, capabilities); + AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { + thinking: Some(Recognized::Known(ThinkingConfig::adaptive(None))), + output_config: with_default_effort(request.params.output_config, level), + ..request.params + }, + ..request + } +} + +fn translate_adaptive_effort_for_non_adaptive_model( + request: AnthropicMessagesRequest, + context: &ThinkingContext, +) -> Result { + let capabilities = &context.capabilities; + if capabilities.supports_adaptive_thinking { + return Ok(request); + } + let effort = known_effort(&request).cloned(); + let adaptive_thinking = matches!(known_thinking(&request), Some(ThinkingConfig::Adaptive(_))); + if effort.is_none() && !adaptive_thinking { + return Ok(request); + } + let level_accepted = match &effort { + Some(Recognized::Known(level)) => capabilities.accepts_effort(*level), + _ => true, + }; + if capabilities.supports_effort_param() && (!adaptive_thinking || level_accepted) { + return Ok(AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { + thinking: if adaptive_thinking { + None + } else { + request.params.thinking + }, + ..request.params + }, + ..request + }); + } + let budget = if capabilities.supports_reasoning { + context + .budgets + .for_effort(legacy_reasoning_effort(effort.as_ref())?) + } else { + None + }; + Ok(AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { + thinking: budget + .and_then(|budget| fit_budget_to_max_tokens(budget, request.params.max_tokens)) + .map(|budget| Recognized::Known(ThinkingConfig::enabled(budget))), + output_config: without_effort(request.params.output_config), + ..request.params + }, + ..request + }) +} + +fn drop_incompatible_temperature_for_thinking( + request: AnthropicMessagesRequest, + context: &ThinkingContext, +) -> AnthropicMessagesRequest { + if context.capabilities.supports_adaptive_thinking { + return request; + } + let pinned = request + .params + .temperature + .is_some_and(|temperature| temperature != 1.0); + let thinking_enabled = matches!(known_thinking(&request), Some(ThinkingConfig::Enabled(_))); + let effort_enabled = known_effort(&request).is_some(); + if !pinned || !(thinking_enabled || effort_enabled) { + return request; + } + AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { + temperature: None, + ..request.params + }, + ..request + } +} + +pub fn translate_thinking( + request: AnthropicMessagesRequest, + context: &ThinkingContext, +) -> Result { + let request = translate_reasoning_effort(request, context)?; + let request = drop_disabled_thinking(request, context); + let request = translate_legacy_thinking_for_adaptive_model(request, context); + let request = translate_adaptive_effort_for_non_adaptive_model(request, context)?; + Ok(drop_incompatible_temperature_for_thinking(request, context)) +} + +#[cfg(test)] +mod tests { + use rstest::{fixture, rstest}; + + use super::*; + use crate::anthropic::common_utils::SupportedEffortTiers; + + const EFFORT_CHOICES: &str = "'none', 'minimal', 'low', 'medium', 'high', 'xhigh', 'max'"; + + fn request(fields: Value) -> AnthropicMessagesRequest { + let mut body = serde_json::json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]}); + body.as_object_mut() + .unwrap() + .extend(fields.as_object().unwrap().clone()); + serde_json::from_value(body).unwrap() + } + + fn context(capabilities: AnthropicModelCapabilities) -> ThinkingContext { + ThinkingContext { + capabilities, + budgets: ThinkingBudgets::default(), + } + } + + fn translate( + capabilities: AnthropicModelCapabilities, + fields: Value, + ) -> Result { + translate_thinking(request(fields), &context(capabilities)) + } + + fn overridden_budgets(overrides: &[(&str, &str)]) -> ThinkingBudgets { + let env = |name: &str| { + overrides + .iter() + .find(|(tier, _)| { + name == format!("DEFAULT_REASONING_EFFORT_{tier}_THINKING_BUDGET") + }) + .map(|(_, value)| value.to_string()) + }; + ThinkingBudgets::from_lookup(&env) + } + + fn claude_code_payload(effort: &str, max_tokens: u64) -> Value { + serde_json::json!({"max_tokens": max_tokens, "thinking": {"type": "adaptive"}, "output_config": {"effort": effort}}) + } + + fn with_temperature(fields: Value, temperature: f64) -> Value { + let mut fields = fields; + fields + .as_object_mut() + .unwrap() + .insert("temperature".to_string(), serde_json::json!(temperature)); + fields + } + + #[fixture] + fn haiku_3_5() -> AnthropicModelCapabilities { + AnthropicModelCapabilities::default() + } + + #[fixture] + fn haiku_4_5() -> AnthropicModelCapabilities { + AnthropicModelCapabilities { + supports_reasoning: true, + ..Default::default() + } + } + + #[fixture] + fn opus_4_5() -> AnthropicModelCapabilities { + AnthropicModelCapabilities { + supports_reasoning: true, + supports_output_config: true, + ..Default::default() + } + } + + #[fixture] + fn sonnet_4_6() -> AnthropicModelCapabilities { + AnthropicModelCapabilities { + supports_reasoning: true, + supports_adaptive_thinking: true, + supports_legacy_thinking: true, + supports_output_config: true, + effort_tiers: SupportedEffortTiers { + max: true, + ..Default::default() + }, + ..Default::default() + } + } + + #[fixture] + fn opus_4_7() -> AnthropicModelCapabilities { + AnthropicModelCapabilities { + supports_reasoning: true, + supports_adaptive_thinking: true, + supports_output_config: true, + effort_tiers: SupportedEffortTiers { + xhigh: true, + max: true, + ..Default::default() + }, + ..Default::default() + } + } + + #[fixture] + fn fable_5_1() -> AnthropicModelCapabilities { + AnthropicModelCapabilities { + thinking_always_on: true, + ..opus_4_7() + } + } + + #[fixture] + fn newfamily_6() -> AnthropicModelCapabilities { + AnthropicModelCapabilities { + supports_reasoning: true, + supports_adaptive_thinking: true, + ..Default::default() + } + } + + #[rstest] + #[case::minimal_maps_to_low(opus_4_7(), "minimal", "low")] + #[case::low(opus_4_7(), "low", "low")] + #[case::medium(opus_4_7(), "medium", "medium")] + #[case::high(opus_4_7(), "high", "high")] + #[case::xhigh_with_xhigh_tier(opus_4_7(), "xhigh", "xhigh")] + #[case::max(opus_4_7(), "max", "max")] + #[case::minimal_maps_to_low_on_4_6(sonnet_4_6(), "minimal", "low")] + #[case::low_on_4_6(sonnet_4_6(), "low", "low")] + #[case::max_without_max_tier_is_allowed_on_adaptive_models(newfamily_6(), "max", "max")] + fn reasoning_effort_on_adaptive_model_becomes_summarized_adaptive_thinking_and_effort( + #[case] capabilities: AnthropicModelCapabilities, + #[case] reasoning_effort: &str, + #[case] expected_effort: &str, + ) { + assert_eq!( + translate( + capabilities, + serde_json::json!({"max_tokens": 1024, "reasoning_effort": reasoning_effort}) + ), + Ok(request(serde_json::json!({ + "max_tokens": 1024, + "thinking": {"type": "adaptive", "display": "summarized"}, + "output_config": {"effort": expected_effort} + }))) + ); + } + + #[rstest] + #[case::adaptive_shape_is_not_dropped_for_small_max_tokens( + opus_4_7(), + serde_json::json!({"max_tokens": 64, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 64, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}}) + )] + #[case::caller_output_config_effort_wins( + opus_4_7(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "low", "output_config": {"effort": "max"}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "max"}}) + )] + #[case::effort_merges_into_caller_output_config( + opus_4_7(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "high", "output_config": {"format": {"type": "json_schema"}}}), + serde_json::json!({ + "max_tokens": 1024, + "thinking": {"type": "adaptive", "display": "summarized"}, + "output_config": {"format": {"type": "json_schema"}, "effort": "high"} + }) + )] + #[case::non_object_output_config_is_replaced( + opus_4_7(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "high", "output_config": "bogus"}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}}) + )] + #[case::caller_thinking_and_output_config_win( + sonnet_4_6(), + serde_json::json!({ + "max_tokens": 16000, + "reasoning_effort": "low", + "thinking": {"type": "enabled", "budget_tokens": 8000}, + "output_config": {"effort": "high"} + }), + serde_json::json!({ + "max_tokens": 16000, + "thinking": {"type": "enabled", "budget_tokens": 8000}, + "output_config": {"effort": "high"} + }) + )] + #[case::caller_legacy_thinking_is_then_translated_while_reasoning_effort_level_stays( + opus_4_7(), + serde_json::json!({"max_tokens": 16000, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), + serde_json::json!({"max_tokens": 16000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "low"}}) + )] + #[case::caller_disabled_thinking_is_kept_then_omitted_on_always_on_model( + fable_5_1(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "high", "thinking": {"type": "disabled"}}), + serde_json::json!({"max_tokens": 1024, "output_config": {"effort": "high"}}) + )] + #[case::non_adaptive_model_gets_no_output_config( + opus_4_5(), + serde_json::json!({"max_tokens": 8192, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + )] + #[case::caller_thinking_wins_on_non_adaptive_model( + opus_4_5(), + serde_json::json!({"max_tokens": 16000, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), + serde_json::json!({"max_tokens": 16000, "thinking": {"type": "enabled", "budget_tokens": 8000}}) + )] + #[case::caller_thinking_survives_when_mapped_budget_cannot_fit( + opus_4_5(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 8000}}) + )] + #[case::missing_max_tokens_leaves_budget_uncapped( + haiku_4_5(), + serde_json::json!({"reasoning_effort": "high"}), + serde_json::json!({"thinking": {"type": "enabled", "budget_tokens": 4096}}) + )] + #[case::budget_below_max_tokens_is_kept( + haiku_4_5(), + serde_json::json!({"max_tokens": 4097, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 4097, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + )] + #[case::budget_equal_to_max_tokens_is_capped( + haiku_4_5(), + serde_json::json!({"max_tokens": 4096, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 4096, "thinking": {"type": "enabled", "budget_tokens": 4095}}) + )] + #[case::budget_above_max_tokens_is_capped( + haiku_4_5(), + serde_json::json!({"max_tokens": 4000, "reasoning_effort": "xhigh"}), + serde_json::json!({"max_tokens": 4000, "thinking": {"type": "enabled", "budget_tokens": 3999}}) + )] + #[case::max_tokens_just_above_min_budget_caps_to_min_budget( + haiku_4_5(), + serde_json::json!({"max_tokens": 1025, "reasoning_effort": "xhigh"}), + serde_json::json!({"max_tokens": 1025, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + )] + #[case::max_tokens_at_min_budget_drops_thinking( + haiku_4_5(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), + serde_json::json!({"max_tokens": 1024}) + )] + #[case::pinned_temperature_is_dropped_after_thinking_is_synthesized( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "reasoning_effort": "low", "temperature": 0}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + )] + #[case::non_string_reasoning_effort_is_ignored( + opus_4_7(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": 3, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "adaptive"}}) + )] + #[case::unrecognized_thinking_is_forwarded( + opus_4_5(), + serde_json::json!({"max_tokens": 8192, "reasoning_effort": "high", "thinking": {"type": "future"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "future"}}) + )] + #[case::caller_display_and_block_binding_survive_on_adaptive_model( + opus_4_7(), + serde_json::json!({ + "max_tokens": 1024, + "reasoning_effort": "high", + "thinking": {"type": "adaptive", "display": "omitted", "block_binding": {"prefix_mismatch_behavior": "drop_block"}} + }), + serde_json::json!({ + "max_tokens": 1024, + "thinking": {"type": "adaptive", "display": "omitted", "block_binding": {"prefix_mismatch_behavior": "drop_block"}}, + "output_config": {"effort": "high"} + }) + )] + fn reasoning_effort_is_translated( + #[case] capabilities: AnthropicModelCapabilities, + #[case] input: Value, + #[case] expected: Value, + ) { + assert_eq!(translate(capabilities, input), Ok(request(expected))); + } + + #[rstest] + #[case::minimal_floors_at_min_budget("minimal", 1024)] + #[case::low("low", 1024)] + #[case::medium("medium", 2048)] + #[case::high("high", 4096)] + #[case::xhigh("xhigh", 8192)] + #[case::max("max", 16384)] + fn reasoning_effort_on_non_adaptive_model_uses_the_tier_budget( + haiku_4_5: AnthropicModelCapabilities, + #[case] reasoning_effort: &str, + #[case] expected_budget: u64, + ) { + assert_eq!( + translate( + haiku_4_5, + serde_json::json!({"max_tokens": 32000, "reasoning_effort": reasoning_effort}) + ), + Ok(request(serde_json::json!({ + "max_tokens": 32000, + "thinking": {"type": "enabled", "budget_tokens": expected_budget} + }))) + ); + } + + #[rstest] + #[case::adaptive_model(opus_4_7())] + #[case::effort_capable_model(opus_4_5())] + #[case::budget_model(haiku_4_5())] + fn reasoning_effort_none_clears_thinking_and_output_config( + #[case] capabilities: AnthropicModelCapabilities, + ) { + assert_eq!( + translate( + capabilities, + serde_json::json!({ + "max_tokens": 1024, + "reasoning_effort": "none", + "thinking": {"type": "adaptive"}, + "output_config": {"effort": "high"} + }) + ), + Ok(request(serde_json::json!({"max_tokens": 1024}))) + ); + } + + #[rstest] + #[case::bogus_on_budget_model( + opus_4_5(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "bogus"}), + format!("Unmapped reasoning effort: 'bogus'. Must be one of: {EFFORT_CHOICES}.") + )] + #[case::disabled_on_budget_model( + haiku_4_5(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "disabled"}), + format!("Unmapped reasoning effort: 'disabled'. Must be one of: {EFFORT_CHOICES}.") + )] + #[case::empty_on_budget_model( + haiku_4_5(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": ""}), + format!("Unmapped reasoning effort: ''. Must be one of: {EFFORT_CHOICES}.") + )] + #[case::invalid_on_adaptive_model( + opus_4_7(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "invalid"}), + format!("Unmapped reasoning effort: 'invalid'. Must be one of: {EFFORT_CHOICES}.") + )] + #[case::disabled_on_adaptive_model( + opus_4_7(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "disabled"}), + format!("Unmapped reasoning effort: 'disabled'. Must be one of: {EFFORT_CHOICES}.") + )] + #[case::empty_on_adaptive_model( + opus_4_7(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": ""}), + format!("Unmapped reasoning effort: ''. Must be one of: {EFFORT_CHOICES}.") + )] + #[case::xhigh_without_xhigh_tier_on_4_6( + sonnet_4_6(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), + "effort='xhigh' is not supported by this model. Got model: claude".to_string() + )] + #[case::xhigh_without_xhigh_tier_on_unmapped_adaptive_model( + newfamily_6(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), + "effort='xhigh' is not supported by this model. Got model: claude".to_string() + )] + #[case::unrecognized_adaptive_effort_on_budget_model( + haiku_4_5(), + claude_code_payload("turbo", 8192), + format!("Unmapped reasoning effort: 'turbo'. Must be one of: {EFFORT_CHOICES}.") + )] + #[case::unrecognized_output_config_effort_on_budget_model( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {"effort": 5}}), + format!("Unmapped reasoning effort: 5. Must be one of: {EFFORT_CHOICES}.") + )] + #[case::quote_in_effort_is_reprd_like_python( + haiku_4_5(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "it's"}), + format!("Unmapped reasoning effort: \"it's\". Must be one of: {EFFORT_CHOICES}.") + )] + fn unsupported_effort_is_a_request_error( + #[case] capabilities: AnthropicModelCapabilities, + #[case] input: Value, + #[case] expected_message: String, + ) { + assert_eq!( + translate(capabilities, input), + Err(Error::InvalidRequest(expected_message)) + ); + } + + #[rstest] + #[case::omitted_on_always_on_model(fable_5_1(), serde_json::json!({"type": "disabled"}), None)] + #[case::kept_on_adaptive_model(opus_4_7(), serde_json::json!({"type": "disabled"}), Some(serde_json::json!({"type": "disabled"})))] + #[case::kept_on_budget_model(haiku_4_5(), serde_json::json!({"type": "disabled"}), Some(serde_json::json!({"type": "disabled"})))] + #[case::adaptive_kept_on_always_on_model( + fable_5_1(), + serde_json::json!({"type": "adaptive"}), + Some(serde_json::json!({"type": "adaptive"})) + )] + fn disabled_thinking_is_omitted_only_for_always_on_models( + #[case] capabilities: AnthropicModelCapabilities, + #[case] thinking: Value, + #[case] expected_thinking: Option, + ) { + let expected = match expected_thinking { + Some(thinking) => serde_json::json!({"max_tokens": 64, "thinking": thinking}), + None => serde_json::json!({"max_tokens": 64}), + }; + assert_eq!( + translate( + capabilities, + serde_json::json!({"max_tokens": 64, "thinking": thinking}) + ), + Ok(request(expected)) + ); + } + + #[rstest] + #[case::far_above_xhigh_budget(opus_4_7(), serde_json::json!(16384), "xhigh")] + #[case::at_xhigh_budget(opus_4_7(), serde_json::json!(8192), "xhigh")] + #[case::below_xhigh_budget(opus_4_7(), serde_json::json!(8191), "high")] + #[case::xhigh_budget_without_xhigh_tier(newfamily_6(), serde_json::json!(8192), "high")] + #[case::large_budget_without_xhigh_tier(newfamily_6(), serde_json::json!(31999), "high")] + #[case::at_high_budget(opus_4_7(), serde_json::json!(4096), "high")] + #[case::below_high_budget(opus_4_7(), serde_json::json!(4095), "medium")] + #[case::at_medium_budget(opus_4_7(), serde_json::json!(2048), "medium")] + #[case::below_medium_budget(opus_4_7(), serde_json::json!(2047), "low")] + #[case::tiny_budget(opus_4_7(), serde_json::json!(1), "low")] + #[case::missing_budget(opus_4_7(), Value::Null, "low")] + #[case::always_on_model(fable_5_1(), serde_json::json!(24000), "xhigh")] + fn legacy_thinking_is_bucketed_into_adaptive_effort_on_adaptive_only_models( + #[case] capabilities: AnthropicModelCapabilities, + #[case] budget_tokens: Value, + #[case] expected_effort: &str, + ) { + let thinking = match budget_tokens { + Value::Null => serde_json::json!({"type": "enabled"}), + budget_tokens => serde_json::json!({"type": "enabled", "budget_tokens": budget_tokens}), + }; + assert_eq!( + translate( + capabilities, + serde_json::json!({"max_tokens": 1024, "thinking": thinking}) + ), + Ok(request(serde_json::json!({ + "max_tokens": 1024, + "thinking": {"type": "adaptive"}, + "output_config": {"effort": expected_effort} + }))) + ); + } + + #[rstest] + #[case::verbatim_on_model_accepting_legacy_thinking( + sonnet_4_6(), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}) + )] + #[case::verbatim_with_explicit_output_config_on_model_accepting_legacy_thinking( + sonnet_4_6(), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}, "output_config": {"effort": "low"}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}, "output_config": {"effort": "low"}}) + )] + #[case::verbatim_on_non_adaptive_model( + opus_4_5(), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}) + )] + #[case::caller_output_config_effort_wins( + opus_4_7(), + serde_json::json!({ + "max_tokens": 32000, + "thinking": {"type": "enabled", "budget_tokens": 31999}, + "output_config": {"effort": "low", "format": {"type": "json_schema"}} + }), + serde_json::json!({ + "max_tokens": 32000, + "thinking": {"type": "adaptive"}, + "output_config": {"effort": "low", "format": {"type": "json_schema"}} + }) + )] + #[case::effort_merges_into_caller_output_config( + opus_4_7(), + serde_json::json!({ + "max_tokens": 32000, + "thinking": {"type": "enabled", "budget_tokens": 4096}, + "output_config": {"format": {"type": "json_schema"}} + }), + serde_json::json!({ + "max_tokens": 32000, + "thinking": {"type": "adaptive"}, + "output_config": {"effort": "high", "format": {"type": "json_schema"}} + }) + )] + #[case::adaptive_thinking_is_left_alone( + opus_4_7(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive", "display": "summarized"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive", "display": "summarized"}}) + )] + fn legacy_thinking_on_adaptive_capable_models( + #[case] capabilities: AnthropicModelCapabilities, + #[case] input: Value, + #[case] expected: Value, + ) { + assert_eq!(translate(capabilities, input), Ok(request(expected))); + } + + #[rstest] + #[case::bare_adaptive_becomes_medium_budget_on_budget_model( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + )] + #[case::medium_effort_becomes_medium_budget_on_budget_model( + haiku_4_5(), + claude_code_payload("medium", 8192), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + )] + #[case::empty_effort_becomes_medium_budget_on_budget_model( + haiku_4_5(), + claude_code_payload("", 8192), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + )] + #[case::high_effort_becomes_high_budget_on_budget_model( + haiku_4_5(), + claude_code_payload("high", 8192), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + )] + #[case::effort_only_becomes_budget_on_budget_model( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + )] + #[case::effort_replaces_caller_legacy_budget_on_budget_model( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 3000}, "output_config": {"effort": "high"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + )] + #[case::residual_output_config_survives_effort_translation( + haiku_4_5(), + serde_json::json!({ + "max_tokens": 8192, + "thinking": {"type": "adaptive"}, + "output_config": {"effort": "medium", "format": {"type": "json_schema"}} + }), + serde_json::json!({ + "max_tokens": 8192, + "thinking": {"type": "enabled", "budget_tokens": 2048}, + "output_config": {"format": {"type": "json_schema"}} + }) + )] + #[case::effortless_output_config_is_kept( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {"format": {"type": "json_schema"}}}), + serde_json::json!({ + "max_tokens": 8192, + "thinking": {"type": "enabled", "budget_tokens": 2048}, + "output_config": {"format": {"type": "json_schema"}} + }) + )] + #[case::empty_output_config_is_kept( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}, "output_config": {}}) + )] + #[case::missing_max_tokens_leaves_budget_uncapped( + haiku_4_5(), + serde_json::json!({"thinking": {"type": "adaptive"}}), + serde_json::json!({"thinking": {"type": "enabled", "budget_tokens": 2048}}) + )] + #[case::budget_is_capped_below_max_tokens( + haiku_4_5(), + claude_code_payload("high", 3000), + serde_json::json!({"max_tokens": 3000, "thinking": {"type": "enabled", "budget_tokens": 2999}}) + )] + #[case::max_tokens_just_above_min_budget_caps_to_min_budget( + haiku_4_5(), + claude_code_payload("medium", 1025), + serde_json::json!({"max_tokens": 1025, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + )] + #[case::max_tokens_at_min_budget_drops_thinking_and_effort( + haiku_4_5(), + claude_code_payload("medium", 1024), + serde_json::json!({"max_tokens": 1024}) + )] + #[case::max_tokens_below_min_budget_drops_thinking_and_effort( + haiku_4_5(), + claude_code_payload("medium", 512), + serde_json::json!({"max_tokens": 512}) + )] + #[case::bare_adaptive_is_dropped_on_non_reasoning_model( + haiku_3_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 8192}) + )] + #[case::adaptive_and_effort_are_dropped_on_non_reasoning_model( + haiku_3_5(), + claude_code_payload("medium", 8192), + serde_json::json!({"max_tokens": 8192}) + )] + #[case::effort_only_is_dropped_on_non_reasoning_model( + haiku_3_5(), + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high", "format": {"type": "json_schema"}}}), + serde_json::json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}) + )] + #[case::supported_effort_is_kept_and_adaptive_thinking_dropped_on_effort_model( + opus_4_5(), + claude_code_payload("medium", 8192), + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "medium"}}) + )] + #[case::bare_adaptive_is_dropped_on_effort_model( + opus_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 8192}) + )] + #[case::effort_only_is_left_alone_on_effort_model( + opus_4_5(), + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}}) + )] + #[case::unsupported_effort_only_is_left_for_provider_normalization( + opus_4_5(), + serde_json::json!({"max_tokens": 4096, "output_config": {"effort": "xhigh"}}), + serde_json::json!({"max_tokens": 4096, "output_config": {"effort": "xhigh"}}) + )] + #[case::legacy_thinking_is_kept_beside_native_effort_on_effort_model( + opus_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}, "output_config": {"effort": "high"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}, "output_config": {"effort": "high"}}) + )] + #[case::unsupported_xhigh_with_adaptive_thinking_falls_back_to_budget( + opus_4_5(), + claude_code_payload("xhigh", 64000), + serde_json::json!({"max_tokens": 64000, "thinking": {"type": "enabled", "budget_tokens": 8192}}) + )] + #[case::unsupported_max_with_adaptive_thinking_falls_back_to_budget( + opus_4_5(), + claude_code_payload("max", 64000), + serde_json::json!({"max_tokens": 64000, "thinking": {"type": "enabled", "budget_tokens": 16384}}) + )] + #[case::bare_adaptive_is_native_on_4_6( + sonnet_4_6(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}) + )] + #[case::adaptive_payload_is_native_on_4_6( + sonnet_4_6(), + claude_code_payload("high", 8192), + claude_code_payload("high", 8192) + )] + #[case::request_without_adaptive_interface_is_left_alone( + haiku_4_5(), + serde_json::json!({"max_tokens": 1024}), + serde_json::json!({"max_tokens": 1024}) + )] + #[case::falsy_non_string_effort_becomes_medium_budget_on_budget_model( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {"effort": 0}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + )] + #[case::minimal_effort_becomes_floored_minimal_budget_on_budget_model( + haiku_4_5(), + claude_code_payload("minimal", 8192), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + )] + #[case::none_effort_drops_thinking_on_budget_model( + haiku_4_5(), + claude_code_payload("none", 8192), + serde_json::json!({"max_tokens": 8192}) + )] + #[case::unrecognized_effort_is_native_on_effort_model( + opus_4_5(), + claude_code_payload("turbo", 8192), + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "turbo"}}) + )] + #[case::task_budget_survives_effort_translation( + haiku_4_5(), + serde_json::json!({ + "max_tokens": 8192, + "thinking": {"type": "adaptive"}, + "output_config": {"effort": "high", "task_budget": {"type": "tokens", "total": 4096}} + }), + serde_json::json!({ + "max_tokens": 8192, + "thinking": {"type": "enabled", "budget_tokens": 4096}, + "output_config": {"task_budget": {"type": "tokens", "total": 4096}} + }) + )] + fn adaptive_interface_is_reshaped_for_non_adaptive_models( + #[case] capabilities: AnthropicModelCapabilities, + #[case] input: Value, + #[case] expected: Value, + ) { + assert_eq!(translate(capabilities, input), Ok(request(expected))); + } + + #[rstest] + #[case::adaptive_downgraded_to_enabled_thinking( + haiku_4_5(), + claude_code_payload("medium", 8192), + 0.0, + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + )] + #[case::bare_adaptive_downgraded_to_enabled_thinking( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), + 0.0, + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + )] + #[case::reasoning_effort_synthesized_enabled_thinking( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "reasoning_effort": "high"}), + 0.2, + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + )] + #[case::above_one_with_enabled_thinking( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}), + 1.5, + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + )] + #[case::native_effort_kept_on_effort_model( + opus_4_5(), + claude_code_payload("medium", 8192), + 0.0, + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "medium"}}) + )] + #[case::effort_only_on_effort_model( + opus_4_5(), + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), + 0.0, + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}}) + )] + fn pinned_temperature_is_dropped_when_thinking_or_effort_survives_on_non_adaptive_model( + #[case] capabilities: AnthropicModelCapabilities, + #[case] input: Value, + #[case] temperature: f64, + #[case] expected: Value, + ) { + assert_eq!( + translate(capabilities, with_temperature(input, temperature)), + Ok(request(expected)) + ); + } + + #[rstest] + #[case::temperature_one_with_enabled_thinking( + haiku_4_5(), + claude_code_payload("medium", 8192), + 1.0, + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + )] + #[case::thinking_dropped_for_small_max_tokens( + haiku_4_5(), + claude_code_payload("medium", 512), + 0.0, + serde_json::json!({"max_tokens": 512}) + )] + #[case::thinking_dropped_on_non_reasoning_model( + haiku_3_5(), + claude_code_payload("medium", 8192), + 0.0, + serde_json::json!({"max_tokens": 8192}) + )] + #[case::disabled_thinking( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "disabled"}}), + 0.0, + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "disabled"}}) + )] + #[case::no_thinking(haiku_4_5(), serde_json::json!({"max_tokens": 8192}), 0.0, serde_json::json!({"max_tokens": 8192}))] + #[case::output_config_without_effort( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}), + 0.0, + serde_json::json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}) + )] + #[case::adaptive_model( + opus_4_7(), + claude_code_payload("medium", 8192), + 0.0, + claude_code_payload("medium", 8192) + )] + #[case::legacy_thinking_on_adaptive_model( + sonnet_4_6(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}), + 0.0, + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + )] + fn temperature_is_kept( + #[case] capabilities: AnthropicModelCapabilities, + #[case] input: Value, + #[case] temperature: f64, + #[case] expected: Value, + ) { + assert_eq!( + translate(capabilities, with_temperature(input, temperature)), + Ok(request(with_temperature(expected, temperature))) + ); + } + + #[rstest] + #[case::minimal("MINIMAL", ThinkingBudgets { minimal: 5000, ..ThinkingBudgets::default() })] + #[case::low("LOW", ThinkingBudgets { low: 5000, ..ThinkingBudgets::default() })] + #[case::medium("MEDIUM", ThinkingBudgets { medium: 5000, ..ThinkingBudgets::default() })] + #[case::high("HIGH", ThinkingBudgets { high: 5000, ..ThinkingBudgets::default() })] + #[case::xhigh("XHIGH", ThinkingBudgets { xhigh: 5000, ..ThinkingBudgets::default() })] + #[case::max("MAX", ThinkingBudgets { max: 5000, ..ThinkingBudgets::default() })] + fn each_tier_budget_reads_only_its_own_environment_override( + #[case] tier: &str, + #[case] expected: ThinkingBudgets, + ) { + assert_eq!(overridden_budgets(&[(tier, "5000")]), expected); + } + + #[rstest] + #[case::whitespace_is_trimmed(" 6000 ", 6000)] + #[case::unparseable_value_keeps_default("lots", 4096)] + fn environment_override_parsing(#[case] raw: &str, #[case] expected_high: u64) { + assert_eq!( + overridden_budgets(&[("HIGH", raw)]), + ThinkingBudgets { + high: expected_high, + ..ThinkingBudgets::default() + } + ); + } + + #[rstest] + #[case::reasoning_effort_uses_overridden_budget( + &[("HIGH", "6000")], + haiku_4_5(), + serde_json::json!({"max_tokens": 32000, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 6000}}) + )] + #[case::minimal_override_below_min_budget_is_floored( + &[("MINIMAL", "512")], + haiku_4_5(), + serde_json::json!({"max_tokens": 32000, "reasoning_effort": "minimal"}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + )] + #[case::minimal_override_above_min_budget_is_used( + &[("MINIMAL", "2000")], + haiku_4_5(), + serde_json::json!({"max_tokens": 32000, "reasoning_effort": "minimal"}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 2000}}) + )] + #[case::adaptive_fallback_uses_overridden_medium_budget( + &[("MEDIUM", "3000")], + haiku_4_5(), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 3000}}) + )] + #[case::legacy_bucket_below_overridden_high_budget( + &[("HIGH", "6000")], + opus_4_7(), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 5999}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}}) + )] + #[case::legacy_bucket_at_overridden_high_budget( + &[("HIGH", "6000")], + opus_4_7(), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 6000}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}) + )] + #[case::legacy_bucket_below_overridden_xhigh_budget( + &[("XHIGH", "20000")], + opus_4_7(), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 19999}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}) + )] + #[case::legacy_bucket_at_overridden_medium_budget( + &[("MEDIUM", "3000")], + opus_4_7(), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 3000}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}}) + )] + #[case::legacy_bucket_below_overridden_medium_budget( + &[("MEDIUM", "3000")], + opus_4_7(), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 2999}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "low"}}) + )] + fn translation_honors_budget_overrides( + #[case] overrides: &[(&str, &str)], + #[case] capabilities: AnthropicModelCapabilities, + #[case] input: Value, + #[case] expected: Value, + ) { + let context = ThinkingContext { + capabilities, + budgets: overridden_budgets(overrides), + }; + assert_eq!( + translate_thinking(request(input), &context), + Ok(request(expected)) + ); + } +} diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/transformation.rs b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs similarity index 54% rename from litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/transformation.rs rename to litellm-rust/crates/llms/src/anthropic/messages/transformation.rs index 59280c04a70..07c2eb46eaa 100644 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs @@ -1,31 +1,35 @@ +use litellm_auth::CredentialPlacement; use litellm_core_utils::settings::{Lookup, ProcessEnvironment}; -use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; +use litellm_types::{ + llms::{ + anthropic::{AnthropicBeta, BetaSet}, + anthropic_messages::anthropic_request::{ + AnthropicMessage, AnthropicMessagesOptionalParams, AnthropicMessagesRequest, + ContextEdit, ContextManagement, Speed, + }, + }, + recognized::Recognized, +}; use serde_json::{Map, Value, json}; -use super::{ - headers::{authenticate, with_feature_betas}, - thinking::{ThinkingBudgets, ThinkingContext, translate_thinking}, -}; +use super::thinking::{ThinkingBudgets, ThinkingContext, translate_thinking}; use crate::{ + Error, anthropic::common_utils::{ - AnthropicModelCapabilities, has_advisor_tool, strip_advisor_blocks, - strip_encrypted_reasoning_blocks, + ANTHROPIC_API_BASE_ENV, ANTHROPIC_API_KEY_ENV, ANTHROPIC_AUTH_TOKEN_ENV, + ANTHROPIC_BASE_URL_ENV, AnthropicModelCapabilities, OauthHandling, complete_anthropic_url, + get_auth_header, has_advisor_tool, has_anthropic_credential, is_tool_search_used, + merge_beta_headers, optionally_handle_anthropic_oauth, requires_native_compaction_beta, + strip_advisor_blocks, strip_encrypted_reasoning_blocks, }, base_llm::{ anthropic_messages::transformation::{ - BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, + BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, ValidatedEnvironment, }, - chat::transformation::Error, + auth::AuthScheme, }, }; -const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY"; -const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN"; -const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE"; -const ANTHROPIC_BASE_URL_ENV: &str = "ANTHROPIC_BASE_URL"; -const DEFAULT_ANTHROPIC_API_BASE: &str = "https://api.anthropic.com"; -const MESSAGES_PATH_SUFFIX: &str = "/v1/messages"; - pub struct AnthropicMessagesConfig; pub const ANTHROPIC_MESSAGES_CONFIG: AnthropicMessagesConfig = AnthropicMessagesConfig; @@ -65,38 +69,31 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig { request: AnthropicMessagesRequest, context: &MessagesTransformContext, ) -> Result { - if request.max_tokens.is_none() { - return Err(Error::InvalidRequest( - "max_tokens is required for Anthropic /v1/messages API".to_string(), - )); + if request.params.max_tokens.is_none() { + return Err(Error::MissingField("max_tokens")); } let request = drop_unsupported_params(request, context)?; let request = translate_thinking(request, &context.thinking)?; let context_management = request + .params .context_management - .as_ref() - .and_then(map_openai_context_management_to_anthropic) - .or_else(|| request.context_management.clone()); - let messages = if has_advisor_tool(request.tools.as_deref()) { + .clone() + .map(map_openai_context_management_to_anthropic); + let messages = if has_advisor_tool(request.params.tools.as_deref()) { request.messages } else { strip_advisor_blocks(request.messages) }; Ok(AnthropicMessagesRequest { messages: strip_encrypted_reasoning_blocks(messages), - context_management, + params: AnthropicMessagesOptionalParams { + context_management, + ..request.params + }, ..request }) } - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_anthropic_api_key(api_key, env_lookup).map_err(Error::from) - } - fn secret_names(&self) -> &'static [&'static str] { &[ ANTHROPIC_API_KEY_ENV, @@ -106,20 +103,108 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig { ] } - fn authenticate( + /// Python's `validate_anthropic_messages_environment` up to the beta merge, which + /// `request_headers` does once the request is final. + fn validate_environment( &self, headers: Headers, api_key: Option<&str>, + _model: &str, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - authenticate(headers, api_key, env_lookup).map_err(Error::from) + ) -> Result { + let headers = match optionally_handle_anthropic_oauth(headers, api_key) { + OauthHandling::Bearer { headers, token } => { + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret: token, + }, + }); + } + OauthHandling::Untouched(headers) => headers, + }; + if has_anthropic_credential(&headers) { + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Forwarded, + }); + } + let auth = get_auth_header(api_key, env_lookup).ok_or(Error::Auth( + litellm_auth::Error::MissingApiKey { + provider: "Anthropic", + environment_variable: ANTHROPIC_API_KEY_ENV, + }, + ))?; + Ok(ValidatedEnvironment { headers, auth }) } fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers { - with_feature_betas(headers, request) + update_headers_with_anthropic_beta(headers, request) } } +fn update_headers_with_anthropic_beta( + headers: Headers, + request: &AnthropicMessagesRequest, +) -> Headers { + merge_beta_headers(headers, feature_betas(request)) +} + +fn feature_betas(request: &AnthropicMessagesRequest) -> BetaSet { + let params = &request.params; + let tools = params.tools.as_deref(); + [ + requires_native_compaction_beta(params.compaction.as_ref(), &request.messages) + .then_some(AnthropicBeta::Compact20260904), + uses_structured_output(params).then_some(AnthropicBeta::StructuredOutputs20251113), + (params.speed == Some(Recognized::Known(Speed::Fast))) + .then_some(AnthropicBeta::FastMode20260201), + messages_carry_output_config(&request.messages) + .then_some(AnthropicBeta::PerTurnControl20260701), + has_advisor_tool(tools).then_some(AnthropicBeta::AdvisorTool20260301), + is_tool_search_used(tools).then_some(AnthropicBeta::AdvancedToolUse20251120), + ] + .into_iter() + .flatten() + .chain(context_management_betas(params.context_management.as_ref())) + .collect() +} + +fn is_compact_edit(edit: &Recognized) -> bool { + matches!(edit, Recognized::Known(ContextEdit::Compact { .. })) +} + +fn context_management_betas( + context_management: Option<&Recognized>, +) -> impl Iterator { + let edits = context_management + .and_then(Recognized::known) + .and_then(|context_management| context_management.edits.as_deref()) + .unwrap_or_default(); + let compact = edits.iter().any(is_compact_edit); + let other = edits.iter().any(|edit| !is_compact_edit(edit)); + compact + .then_some(AnthropicBeta::Compact20260112) + .into_iter() + .chain(other.then_some(AnthropicBeta::ContextManagement20250627)) +} + +fn uses_structured_output(params: &AnthropicMessagesOptionalParams) -> bool { + params.output_format.is_some() + || params + .output_config + .as_ref() + .and_then(Recognized::known) + .is_some_and(|config| config.format.is_some()) +} + +fn messages_carry_output_config(messages: &[AnthropicMessage]) -> bool { + messages + .iter() + .any(|message| message.extra.contains_key("output_config")) +} + fn unsupported_param(model: &str, param: &str, value: &str, hint: &str) -> Error { Error::InvalidRequest(format!( "{model} does not support {param}={value}. {hint}To drop unsupported params, set `litellm.drop_params = True`." @@ -138,17 +223,21 @@ fn drop_unsupported_params( } Err(unsupported_param(&model, param, &value, hint)) }; - let speed = match request.speed.as_deref() { + let params = request.params; + let speed = match ¶ms.speed { Some(speed) if !capabilities.supports_speed => { - reject("speed", format!("'{speed}'"), "")?; + reject("speed", format!("'{}'", speed_text(speed)), "")?; None } - _ => request.speed.clone(), + _ => params.speed.clone(), }; if capabilities.supports_sampling_params { - return Ok(AnthropicMessagesRequest { speed, ..request }); + return Ok(AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { speed, ..params }, + ..request + }); } - let temperature = match request.temperature { + let temperature = match params.temperature { Some(temperature) if temperature != 1.0 => { reject( "temperature", @@ -159,92 +248,74 @@ fn drop_unsupported_params( } temperature => temperature, }; - if let Some(top_p) = request.top_p { + if let Some(top_p) = params.top_p { reject("top_p", json!(top_p).to_string(), "")?; } - if let Some(top_k) = request.top_k { + if let Some(top_k) = params.top_k { reject("top_k", json!(top_k).to_string(), "")?; } Ok(AnthropicMessagesRequest { - speed, - temperature, - top_p: None, - top_k: None, + params: AnthropicMessagesOptionalParams { + speed, + temperature, + top_p: None, + top_k: None, + ..params + }, ..request }) } -pub fn map_openai_context_management_to_anthropic(context_management: &Value) -> Option { - match context_management { - Value::Object(edits) if edits.contains_key("edits") => Some(context_management.clone()), - Value::Array(entries) => { - let edits: Vec = entries - .iter() - .filter_map(Value::as_object) - .filter(|entry| entry.get("type").and_then(Value::as_str) == Some("compaction")) - .map(|entry| { - let trigger = entry.get("compact_threshold").and_then(Value::as_f64).map( - |threshold| json!({"type": "input_tokens", "value": threshold as i64}), - ); - let passthrough = entry - .iter() - .filter(|(key, _)| !matches!(key.as_str(), "type" | "compact_threshold")) - .map(|(key, value)| (key.clone(), value.clone())); - Value::Object( - [("type".to_string(), json!("compact_20260112"))] - .into_iter() - .chain(trigger.map(|trigger| ("trigger".to_string(), trigger))) - .chain(passthrough) - .collect::>(), - ) - }) - .collect(); - (!edits.is_empty()).then(|| json!({"edits": edits})) - } - _ => None, +fn speed_text(speed: &Recognized) -> String { + match speed { + Recognized::Known(speed) => speed.as_str().to_string(), + Recognized::Unrecognized(Value::String(text)) => text.clone(), + Recognized::Unrecognized(other) => other.to_string(), } } -pub fn non_empty(value: Option<&str>) -> Option<&str> { - value.map(str::trim).filter(|value| !value.is_empty()) -} - -pub fn resolve_anthropic_api_key( - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - non_empty(api_key) - .map(str::to_string) - .or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty())) - .ok_or(litellm_auth::Error::MissingApiKey { - provider: "Anthropic", - environment_variable: ANTHROPIC_API_KEY_ENV, - }) -} - -pub fn complete_anthropic_url( - api_base: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> String { - let api_base = resolve_anthropic_api_base(api_base, env_lookup); - - let api_base = api_base.trim_end_matches('/'); - if api_base.ends_with(MESSAGES_PATH_SUFFIX) { - return api_base.to_string(); +fn compact_edit_from_openai(entry: &Map) -> Option { + if entry.get("type").and_then(Value::as_str) != Some("compaction") { + return None; } - format!("{api_base}{MESSAGES_PATH_SUFFIX}") + let trigger = entry + .get("compact_threshold") + .and_then(Value::as_f64) + .map(|threshold| json!({"type": "input_tokens", "value": threshold as i64})); + let passthrough = entry + .iter() + .filter(|(key, _)| !matches!(key.as_str(), "type" | "compact_threshold")) + .map(|(key, value)| (key.clone(), value.clone())); + Some(ContextEdit::Compact { + extra: trigger + .map(|trigger| ("trigger".to_string(), trigger)) + .into_iter() + .chain(passthrough) + .collect(), + }) } -pub fn resolve_anthropic_api_base( - api_base: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> String { - let env = |name: &str| env_lookup(name).filter(|value| !value.trim().is_empty()); - non_empty(api_base) - .map(str::to_string) - .or_else(|| env(ANTHROPIC_API_BASE_ENV)) - .or_else(|| env(ANTHROPIC_BASE_URL_ENV)) - .unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string()) +/// An OpenAI-style `context_management` list becomes Anthropic `edits` when it holds +/// compaction entries. Anything else, native edits included, is sent as it came. +pub fn map_openai_context_management_to_anthropic( + context_management: Recognized, +) -> Recognized { + let Recognized::Unrecognized(Value::Array(entries)) = &context_management else { + return context_management; + }; + let edits: Vec> = entries + .iter() + .filter_map(Value::as_object) + .filter_map(compact_edit_from_openai) + .map(Recognized::Known) + .collect(); + if edits.is_empty() { + return context_management; + } + Recognized::Known(ContextManagement { + edits: Some(edits), + extra: Map::new(), + }) } #[cfg(test)] @@ -254,17 +325,14 @@ mod tests { use rstest::{fixture, rstest}; use super::*; - use crate::anthropic::common_utils::{ENCRYPTED_REASONING_SIGNATURE_PREFIX, beta}; + use crate::anthropic::common_utils::ENCRYPTED_REASONING_SIGNATURE_PREFIX; type Env = &'static [(&'static str, &'static str)]; - const BOTH_BASE_ENVS: Env = &[ - (ANTHROPIC_API_BASE_ENV, "https://api-base.example.com"), - (ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com"), - ]; - const API_KEY_ENV: Env = &[(ANTHROPIC_API_KEY_ENV, "sk-env")]; - const MISSING_API_KEY: &str = - "Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"; + const OAUTH_TOKEN: &str = "sk-ant-oat01-token"; + const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token"; + const OAUTH_BETA: &str = "oauth-2025-04-20"; + const BROWSER_ACCESS: (&str, &str) = ("anthropic-dangerous-direct-browser-access", "true"); const LOW_BUDGET_ENV: &str = "DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET"; const PROCESS_ENV_PROBE: &str = "LITELLM_MESSAGES_TRANSFORM_CONTEXT_PROBE"; @@ -369,7 +437,7 @@ mod tests { fn missing_max_tokens_is_rejected(#[case] fields: Value, unmapped: AnthropicModelCapabilities) { assert_eq!( transform(fields, unmapped, false), - invalid("max_tokens is required for Anthropic /v1/messages API") + Err(Error::MissingField("max_tokens")) ); } @@ -561,17 +629,19 @@ mod tests { #[case::empty_list(json!([]), None)] #[case::anthropic_edits_pass_through( json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}), - Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]})) + None )] #[case::object_without_edits(json!({"type": "compaction"}), None)] #[case::scalar(json!("compaction"), None)] fn openai_context_management_maps_to_anthropic_edits( #[case] context_management: Value, - #[case] expected: Option, + #[case] mapped: Option, ) { + let parsed: Recognized = + serde_json::from_value(context_management.clone()).unwrap(); assert_eq!( - map_openai_context_management_to_anthropic(&context_management), - expected + serde_json::to_value(map_openai_context_management_to_anthropic(parsed)).unwrap(), + mapped.unwrap_or(context_management) ); } @@ -715,47 +785,6 @@ mod tests { ); } - #[rstest] - #[case::public_endpoint_by_default(None, &[], "https://api.anthropic.com")] - #[case::explicit_api_base_beats_env( - Some("https://explicit.example.com"), - BOTH_BASE_ENVS, - "https://explicit.example.com" - )] - #[case::explicit_api_base_is_trimmed( - Some(" https://explicit.example.com "), - &[], - "https://explicit.example.com" - )] - #[case::blank_api_base_falls_back_to_env( - Some(" "), - BOTH_BASE_ENVS, - "https://api-base.example.com" - )] - #[case::api_base_env_beats_base_url_env(None, BOTH_BASE_ENVS, "https://api-base.example.com")] - #[case::base_url_env_without_api_base_env( - None, - &[(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")], - "https://base-url.example.com" - )] - #[case::blank_api_base_env_falls_back_to_base_url_env( - None, - &[(ANTHROPIC_API_BASE_ENV, " \t "), (ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")], - "https://base-url.example.com" - )] - #[case::blank_envs_fall_back_to_public_endpoint( - None, - &[(ANTHROPIC_API_BASE_ENV, ""), (ANTHROPIC_BASE_URL_ENV, " ")], - "https://api.anthropic.com" - )] - fn api_base_resolution( - #[case] api_base: Option<&str>, - #[case] vars: Env, - #[case] expected: &str, - ) { - assert_eq!(resolve_anthropic_api_base(api_base, &env(vars)), expected); - } - #[rstest] #[case::public_endpoint(None, &[], "https://api.anthropic.com/v1/messages")] #[case::base_url_env( @@ -786,78 +815,278 @@ mod tests { ); } + fn betas(values: &[&str]) -> BetaSet { + values.join(",").parse().unwrap() + } + + fn validated( + forwarded: &[(&str, &str)], + api_key: Option<&str>, + vars: Env, + ) -> Result { + ANTHROPIC_MESSAGES_CONFIG.validate_environment( + headers(forwarded), + api_key, + "claude", + &env(vars), + ) + } + + fn credential(auth: &AuthScheme) -> Option<(&'static str, &str)> { + match auth { + AuthScheme::Credential { placement, secret } => { + Some((placement.header_name(), secret.expose())) + } + AuthScheme::Forwarded => None, + other => panic!("unexpected auth scheme {other:?}"), + } + } + #[rstest] - #[case::param_beats_env(Some("sk-param"), API_KEY_ENV, Ok("sk-param"))] - #[case::param_is_trimmed(Some(" sk-param "), &[], Ok("sk-param"))] - #[case::blank_param_falls_back_to_env(Some(" "), API_KEY_ENV, Ok("sk-env"))] - #[case::env_without_param(None, API_KEY_ENV, Ok("sk-env"))] - #[case::blank_env_is_missing(None, &[(ANTHROPIC_API_KEY_ENV, " ")], Err(MISSING_API_KEY))] - #[case::nothing_is_missing(None, &[], Err(MISSING_API_KEY))] - fn api_key_resolution( + #[case::forwarded_oauth_bearer( + &[("anthropic-version", "2023-06-01"), ("X-Api-Key", "sk-caller"), ("Authorization", OAUTH_BEARER)], + Some("sk-deployment"), + &[("ANTHROPIC_API_KEY", "sk-env")], + &[("anthropic-version", "2023-06-01"), ("anthropic-beta", OAUTH_BETA), BROWSER_ACCESS], + Some(("Authorization", OAUTH_TOKEN)), + )] + #[case::oauth_api_key( + &[("x-api-key", OAUTH_TOKEN), ("anthropic-beta", "web-search-2025-03-05")], + Some(OAUTH_TOKEN), + &[], + &[("anthropic-beta", "oauth-2025-04-20,web-search-2025-03-05"), BROWSER_ACCESS], + Some(("Authorization", OAUTH_TOKEN)), + )] + #[case::forwarded_x_api_key_is_kept_over_the_deployment_key( + &[("X-API-KEY", "caller-key")], + Some("sk-other"), + &[("ANTHROPIC_API_KEY", "sk-env")], + &[("X-API-KEY", "caller-key")], + None, + )] + #[case::forwarded_non_oauth_bearer_is_kept( + &[("Authorization", "Bearer some-proxy-token")], + Some("sk-ant-api03-regular"), + &[], + &[("Authorization", "Bearer some-proxy-token")], + None, + )] + #[case::oauth_token_without_the_bearer_scheme_is_kept( + &[("authorization", OAUTH_TOKEN)], + None, + &[], + &[("authorization", OAUTH_TOKEN)], + None, + )] + #[case::api_key_param( + &[("anthropic-beta", "web-search-2025-03-05")], + Some("sk-param"), + &[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")], + &[("anthropic-beta", "web-search-2025-03-05")], + Some(("x-api-key", "sk-param")), + )] + #[case::env_key_when_the_param_is_blank( + &[], + Some(" "), + &[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")], + &[], + Some(("x-api-key", "sk-env")), + )] + #[case::auth_token_when_no_key_is_set( + &[], + None, + &[("ANTHROPIC_API_KEY", " \t"), ("ANTHROPIC_AUTH_TOKEN", "env-token")], + &[], + Some(("Authorization", "env-token")), + )] + #[case::oauth_env_key_as_a_bearer( + &[], + None, + &[("ANTHROPIC_API_KEY", "sk-ant-oat01-env")], + &[], + Some(("Authorization", "sk-ant-oat01-env")), + )] + fn validate_environment_shapes_the_headers_and_names_the_credential( + #[case] forwarded: &[(&str, &str)], #[case] api_key: Option<&str>, #[case] vars: Env, - #[case] expected: Result<&str, &str>, + #[case] expected_headers: &[(&str, &str)], + #[case] expected_credential: Option<(&str, &str)>, ) { - assert_eq!( - resolve_anthropic_api_key(api_key, &env(vars)).map_err(|error| error.to_string()), - expected.map(str::to_string).map_err(str::to_string) - ); - } - - #[test] - fn config_reports_a_missing_key_as_an_auth_error() { - assert_eq!( - ANTHROPIC_MESSAGES_CONFIG.resolve_api_key(None, &no_env), - Err(Error::Auth(litellm_auth::Error::MissingApiKey { - provider: "Anthropic", - environment_variable: ANTHROPIC_API_KEY_ENV, - })) - ); - } - - #[test] - fn config_authenticates_with_the_anthropic_auth_token() { - assert_eq!( - ANTHROPIC_MESSAGES_CONFIG.authenticate( - vec![], - None, - &env(&[("ANTHROPIC_AUTH_TOKEN", "auth-token")]) - ), - Ok(headers(&[("authorization", "Bearer auth-token")])) - ); - } - - #[test] - fn config_requests_the_betas_the_request_features_need() { - assert_eq!( - ANTHROPIC_MESSAGES_CONFIG.request_headers( - headers(&[("x-api-key", "sk")]), - &request(json!({"speed": "fast"})) - ), - headers(&[ - ("x-api-key", "sk"), - ("anthropic-beta", beta::FAST_MODE_2026_02_01) - ]) - ); + let environment = validated(forwarded, api_key, vars).unwrap(); + assert_eq!(environment.headers, headers(expected_headers)); + assert_eq!(credential(&environment.auth), expected_credential); } #[rstest] - #[case::absent(None, None)] - #[case::blank(Some(" \t "), None)] - #[case::padded(Some(" value "), Some("value"))] - fn non_empty_trims_and_drops_blank_values( - #[case] value: Option<&str>, - #[case] expected: Option<&str>, + #[case::no_credentials(&[], None, &[])] + #[case::empty_api_key(&[], Some(""), &[])] + #[case::whitespace_only_env_values(&[], None, &[("ANTHROPIC_API_KEY", " "), ("ANTHROPIC_AUTH_TOKEN", " \t")])] + #[case::unrelated_forwarded_headers(&[("anthropic-beta", "web-search-2025-03-05")], None, &[])] + fn missing_credentials_are_an_auth_error( + #[case] forwarded: &[(&str, &str)], + #[case] api_key: Option<&str>, + #[case] vars: Env, ) { - assert_eq!(non_empty(value), expected); + assert!(matches!( + validated(forwarded, api_key, vars), + Err(Error::Auth(litellm_auth::Error::MissingApiKey { + provider: "Anthropic", + environment_variable: "ANTHROPIC_API_KEY", + })) + )); + } + + #[rstest] + #[case::no_features(json!({}), &[])] + #[case::output_format(json!({"output_format": {"type": "json_schema"}}), &["structured-outputs-2025-11-13"])] + #[case::null_output_format(json!({"output_format": null}), &[])] + #[case::output_config_format( + json!({"output_config": {"format": {"type": "json_schema"}, "effort": "xhigh"}}), + &["structured-outputs-2025-11-13"] + )] + #[case::null_output_config_format(json!({"output_config": {"format": null}}), &[])] + #[case::top_level_output_config_without_format(json!({"output_config": {"effort": "high"}}), &[])] + #[case::fast_speed(json!({"speed": "fast"}), &["fast-mode-2026-02-01"])] + #[case::standard_speed(json!({"speed": "standard"}), &[])] + #[case::unknown_speed(json!({"speed": "turbo"}), &[])] + #[case::compaction_param(json!({"compaction": {"enabled": true}}), &["compact-2026-09-04"])] + #[case::empty_compaction_param(json!({"compaction": {}}), &["compact-2026-09-04"])] + #[case::signed_compaction_block_in_history( + json!({"messages": [ + {"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": "sig"}]}, + {"role": "user", "content": "Continue"}, + ]}), + &["compact-2026-09-04"] + )] + #[case::unsigned_compaction_block_in_history( + json!({"messages": [ + {"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": ""}]}, + {"role": "user", "content": "Continue"}, + ]}), + &[] + )] + #[case::advisor_tool( + json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": "claude-opus-4-6"}]}), + &["advisor-tool-2026-03-01"] + )] + #[case::no_tools(json!({"tools": []}), &[])] + #[case::regex_tool_search( + json!({"tools": [{"type": "tool_search_tool_regex_20251119"}]}), + &["advanced-tool-use-2025-11-20"] + )] + #[case::bm25_tool_search( + json!({"tools": [{"type": "tool_search_tool_bm25_20251119"}]}), + &["advanced-tool-use-2025-11-20"] + )] + #[case::unrelated_server_tool(json!({"tools": [{"type": "web_search_20250305", "name": "web_search"}]}), &[])] + #[case::only_compact_edits( + json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}), + &["compact-2026-01-12"] + )] + #[case::only_other_edits( + json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919", "keep": {"type": "tool_uses", "value": 3}}]}}), + &["context-management-2025-06-27"] + )] + #[case::compact_and_other_edits( + json!({"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_tool_uses_20250919"}]}}), + &["compact-2026-01-12", "context-management-2025-06-27"] + )] + #[case::edit_without_a_type(json!({"context_management": {"edits": [{}]}}), &["context-management-2025-06-27"])] + #[case::unknown_edit_type(json!({"context_management": {"edits": [{"type": "future"}]}}), &["context-management-2025-06-27"])] + #[case::empty_edits(json!({"context_management": {"edits": []}}), &[])] + #[case::context_management_without_edits(json!({"context_management": {}}), &[])] + #[case::unmapped_openai_context_management(json!({"context_management": [{"type": "other"}]}), &[])] + #[case::per_message_output_config( + json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}), + &["per-turn-control-2026-07-01"] + )] + #[case::per_message_null_output_config( + json!({"messages": [{"role": "user", "content": "hi", "output_config": null}]}), + &["per-turn-control-2026-07-01"] + )] + fn feature_betas_follow_the_request(#[case] fields: Value, #[case] expected: &[&str]) { + assert_eq!(feature_betas(&request(fields)), betas(expected)); + } + + #[rstest] + #[case::no_betas(&[("x-api-key", "k"), ("anthropic-version", "2023-06-01")], json!({}), &[("x-api-key", "k"), ("anthropic-version", "2023-06-01")])] + #[case::blank_beta_header(&[("Anthropic-Beta", " , "), ("x-api-key", "k")], json!({}), &[("Anthropic-Beta", " , "), ("x-api-key", "k")])] + #[case::feature_beta_is_appended( + &[("x-api-key", "k")], + json!({"speed": "fast"}), + &[("x-api-key", "k"), ("anthropic-beta", "fast-mode-2026-02-01")], + )] + #[case::existing_betas_are_normalized_without_features( + &[("Anthropic-Beta", "web-search-2025-03-05, interleaved-thinking-2025-05-14 ,web-search-2025-03-05"), ("x-api-key", "k")], + json!({}), + &[("x-api-key", "k"), ("anthropic-beta", "interleaved-thinking-2025-05-14,web-search-2025-03-05")], + )] + #[case::existing_advisor_beta_is_kept_without_an_advisor_tool( + &[("anthropic-beta", "advisor-tool-2026-03-01")], + json!({"tools": []}), + &[("anthropic-beta", "advisor-tool-2026-03-01")], + )] + #[case::feature_already_sent_is_not_duplicated( + &[("anthropic-beta", "fast-mode-2026-02-01")], + json!({"speed": "fast"}), + &[("anthropic-beta", "fast-mode-2026-02-01")], + )] + #[case::differently_cased_beta_header_is_replaced_by_one_sorted_header( + &[("Anthropic-Beta", "interleaved-thinking-2025-05-14")], + json!({"messages": [{"role": "system", "content": "env", "output_config": {"effort": "low"}}]}), + &[("anthropic-beta", "interleaved-thinking-2025-05-14,per-turn-control-2026-07-01")], + )] + #[case::every_beta_header_casing_is_unioned_into_one_header( + &[("anthropic-beta", "interleaved-thinking-2025-05-14"), ("Anthropic-Beta", "web-search-2025-03-05")], + json!({"speed": "fast"}), + &[("anthropic-beta", "fast-mode-2026-02-01,interleaved-thinking-2025-05-14,web-search-2025-03-05")], + )] + #[case::unknown_client_betas_survive_alongside_the_added_one( + &[("anthropic-beta", "claude-code-20250219,interleaved-thinking-2025-05-14,context-management-2025-06-27,per-turn-control-2026-07-01,effort-2025-11-24")], + json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}), + &[("anthropic-beta", "claude-code-20250219,context-management-2025-06-27,effort-2025-11-24,interleaved-thinking-2025-05-14,per-turn-control-2026-07-01")], + )] + fn request_headers_merge_the_feature_betas( + #[case] input: &[(&str, &str)], + #[case] fields: Value, + #[case] expected: &[(&str, &str)], + ) { + assert_eq!( + ANTHROPIC_MESSAGES_CONFIG.request_headers(headers(input), &request(fields)), + headers(expected) + ); } #[test] - fn auth_strategy_and_default_headers_match_anthropic() { + fn every_feature_merges_with_the_oauth_beta_sorted() { + let environment = validated(&[], Some(OAUTH_TOKEN), &[]).unwrap(); + let all_features = request(json!({ + "compaction": {"enabled": true}, + "output_format": {"type": "json_schema"}, + "speed": "fast", + "tools": [{"type": "advisor_20260301"}, {"type": "tool_search_tool_bm25_20251119"}], + "context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_thinking_20251015"}]}, + "messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}], + })); assert_eq!( - ANTHROPIC_MESSAGES_CONFIG.auth_strategy().header_name(), - "x-api-key" + ANTHROPIC_MESSAGES_CONFIG.request_headers(environment.headers, &all_features), + headers(&[ + BROWSER_ACCESS, + ( + "anthropic-beta", + "advanced-tool-use-2025-11-20,advisor-tool-2026-03-01,compact-2026-01-12,compact-2026-09-04,context-management-2025-06-27,fast-mode-2026-02-01,oauth-2025-04-20,per-turn-control-2026-07-01,structured-outputs-2025-11-13" + ), + ]) ); + assert_eq!( + credential(&environment.auth), + Some(("Authorization", OAUTH_TOKEN)) + ); + } + + #[test] + fn default_headers_match_anthropic() { assert_eq!( ANTHROPIC_MESSAGES_CONFIG.default_headers(), &[ @@ -874,7 +1103,7 @@ mod tests { requested.borrow_mut().push(name.to_string()); None }; - let _ = ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record); + let _ = ANTHROPIC_MESSAGES_CONFIG.validate_environment(Vec::new(), None, "claude", &record); let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record); let requested = requested.into_inner(); assert!(!requested.is_empty()); diff --git a/litellm-rust/crates/llms/src/anthropic/mod.rs b/litellm-rust/crates/llms/src/anthropic/mod.rs index 755bc7d1907..a884c146dca 100644 --- a/litellm-rust/crates/llms/src/anthropic/mod.rs +++ b/litellm-rust/crates/llms/src/anthropic/mod.rs @@ -1,7 +1,8 @@ +pub mod common_utils; + pub mod batches; pub mod chat; -pub mod common_utils; pub mod count_tokens; -pub mod experimental_pass_through; +pub mod messages; pub const ANTHROPIC_OAUTH_TOKEN_PREFIX: &str = "sk-ant-oat"; diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs index 2ce1b0da51b..1defec654bd 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs @@ -63,9 +63,14 @@ impl BaseOcrConfig for TextractAnalyzeDocumentConfig { async fn validate_environment( &self, request: &PreparedOcrRequest, - _client: &OcrClient, + client: &OcrClient, ) -> Result { - environment(request, TextractOperation::AnalyzeDocument).await + environment( + &client.auth().aws, + request, + TextractOperation::AnalyzeDocument, + ) + .await } fn get_complete_url( diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs index 8268ad066a1..678104be982 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs @@ -1,5 +1,5 @@ use base64::{Engine, engine::general_purpose::STANDARD}; -use litellm_auth_aws::{SigV4Signer, resolve_aws_region}; +use litellm_auth_aws::{AwsCredentialSource, SigV4Signer, resolve_aws_region}; use litellm_http::outbound::RequestSigner; use serde::{Deserialize, Serialize}; use strum::{EnumString, IntoStaticStr, VariantNames}; @@ -232,6 +232,7 @@ pub(super) fn health_check_document() -> OcrDocument { } pub(super) async fn environment( + auth: &litellm_auth_aws::AwsAuthService, request: &PreparedOcrRequest, operation: TextractOperation, ) -> Result { @@ -244,9 +245,10 @@ pub(super) async fn environment( ) })?; let signer = SigV4Signer::resolve( + auth, region.clone(), TEXTRACT_SERVICE, - &request.optional_params, + AwsCredentialSource::from_params(&request.optional_params, &env_lookup), &env_lookup, ) .await diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs index 6eb195defaa..3b4f8e5a7d9 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs @@ -48,9 +48,14 @@ impl BaseOcrConfig for TextractDetectTextConfig { async fn validate_environment( &self, request: &PreparedOcrRequest, - _client: &OcrClient, + client: &OcrClient, ) -> Result { - environment(request, TextractOperation::DetectDocumentText).await + environment( + &client.auth().aws, + request, + TextractOperation::DetectDocumentText, + ) + .await } fn get_complete_url( diff --git a/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs b/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs index c409f7f687e..0c1434afb5d 100644 --- a/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs @@ -1,26 +1,30 @@ +use litellm_auth::SecretValue; +use litellm_http::request::{has_bearer_auth, has_header}; use litellm_types::llms::anthropic_messages::{ anthropic_request::{ - AnthropicMessage, AnthropicMessagesRequest, ContentBlock, MessageContent, SystemPrompt, + AnthropicMessage, AnthropicMessagesOptionalParams, AnthropicMessagesRequest, ContentBlock, + MessageContent, SystemPrompt, }, anthropic_response::AnthropicMessagesResponse, }; use crate::{ - anthropic::experimental_pass_through::messages::transformation::{ - ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty, + Error, + anthropic::{ + common_utils::{API_KEY_PLACEMENT, MESSAGES_PATH_SUFFIX, non_empty}, + messages::transformation::{ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig}, }, base_llm::{ anthropic_messages::transformation::{ - BaseAnthropicMessagesConfig, Headers, MessagesAuthStrategy, MessagesTransformContext, + BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, ValidatedEnvironment, }, - chat::transformation::Error, + auth::AuthScheme, }, }; const AZURE_API_KEY_ENV: &str = "AZURE_API_KEY"; const AZURE_API_BASE_ENV: &str = "AZURE_API_BASE"; const ANTHROPIC_PATH_SEGMENT: &str = "/anthropic"; -const MESSAGES_PATH_SUFFIX: &str = "/v1/messages"; const SYSTEM_ROLE: &str = "system"; pub struct AzureAnthropicMessagesConfig { @@ -48,7 +52,7 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig { context: &MessagesTransformContext, ) -> Result { let mut request = fold_system_role_messages(request); - if let Some(system) = request.system.as_mut() { + if let Some(system) = request.params.system.as_mut() { strip_scope_from_system(system); } request @@ -68,24 +72,30 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig { .transform_anthropic_messages_response(model, response) } - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_azure_api_key(api_key, env_lookup) - } - fn secret_names(&self) -> &'static [&'static str] { &[AZURE_API_KEY_ENV, AZURE_API_BASE_ENV] } - fn auth_strategy(&self) -> MessagesAuthStrategy { - self.anthropic.auth_strategy() - } - - fn accepts_bearer_auth(&self) -> bool { - true + /// A forwarded `x-api-key` or a non-blank bearer (an Entra ID token) is the credential; + /// otherwise the Azure key goes in `x-api-key`. + fn validate_environment( + &self, + headers: Headers, + api_key: Option<&str>, + _model: &str, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + if has_header(&headers, API_KEY_PLACEMENT.header_name()) || has_bearer_auth(&headers) { + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Forwarded, + }); + } + let auth = AuthScheme::Credential { + placement: API_KEY_PLACEMENT, + secret: SecretValue::new(resolve_azure_api_key(api_key, env_lookup)?), + }; + Ok(ValidatedEnvironment { headers, auth }) } fn default_headers(&self) -> &'static [(&'static str, &'static str)] { @@ -181,7 +191,7 @@ fn fold_system_role_messages(request: AnthropicMessagesRequest) -> AnthropicMess .into_iter() .partition(|msg| msg.role == SYSTEM_ROLE); - let folded_system: Vec = system_into_blocks(request.system) + let folded_system: Vec = system_into_blocks(request.params.system) .into_iter() .chain( system_messages @@ -192,15 +202,21 @@ fn fold_system_role_messages(request: AnthropicMessagesRequest) -> AnthropicMess AnthropicMessagesRequest { messages: chat_messages, - system: (!folded_system.is_empty()).then_some(SystemPrompt::Blocks(folded_system)), + params: AnthropicMessagesOptionalParams { + system: (!folded_system.is_empty()).then_some(SystemPrompt::Blocks(folded_system)), + ..request.params + }, ..request } } #[cfg(test)] mod tests { + use rstest::rstest; use serde_json::json; + use litellm_auth::CredentialPlacement; + use super::*; use crate::anthropic::common_utils::AnthropicModelCapabilities; @@ -292,19 +308,47 @@ mod tests { )); } - #[test] - fn auth_strategy_is_x_api_key() { - assert_eq!( - AZURE_ANTHROPIC_MESSAGES_CONFIG - .auth_strategy() - .header_name(), - "x-api-key" - ); + fn validated(forwarded: &[(&str, &str)], api_key: Option<&str>) -> ValidatedEnvironment { + AZURE_ANTHROPIC_MESSAGES_CONFIG + .validate_environment( + forwarded + .iter() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect(), + api_key, + "claude", + &|_| None, + ) + .unwrap() } #[test] - fn accepts_bearer_auth_for_entra_id() { - assert!(AZURE_ANTHROPIC_MESSAGES_CONFIG.accepts_bearer_auth()); + fn the_azure_key_goes_in_x_api_key() { + assert!(matches!( + validated(&[], Some("sk-azure")).auth, + AuthScheme::Credential { + placement: CredentialPlacement::Header("x-api-key"), + ref secret + } if secret.expose() == "sk-azure" + )); + } + + #[rstest] + #[case::x_api_key(&[("X-Api-Key", "caller")])] + #[case::entra_id_bearer(&[("Authorization", "Bearer eyJ-token")])] + fn a_forwarded_key_or_bearer_is_the_credential(#[case] forwarded: &[(&str, &str)]) { + assert!(matches!( + validated(forwarded, Some("sk-azure")).auth, + AuthScheme::Forwarded + )); + } + + #[test] + fn a_blank_bearer_does_not_count_as_a_credential() { + assert!(matches!( + validated(&[("Authorization", "Bearer ")], Some("sk-azure")).auth, + AuthScheme::Credential { .. } + )); } #[test] @@ -528,7 +572,7 @@ mod tests { assert!(err.is_data()); } - #[rstest::rstest] + #[rstest] #[case::compact_context_management_edit( json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}), &[], @@ -608,7 +652,12 @@ mod tests { requested.borrow_mut().push(name.to_string()); None }; - let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record); + let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.validate_environment( + Vec::new(), + None, + "claude", + &record, + ); let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record); let requested = requested.into_inner(); assert!(!requested.is_empty()); diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/common_utils.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/common_utils.rs index 9c2f3f70b91..f28e5c135b0 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/common_utils.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/common_utils.rs @@ -1,5 +1,3 @@ -use std::sync::OnceLock; - use litellm_auth::{InputSource, Sourced}; use litellm_auth_azure::{AzureAuthInputs, AzureAuthService}; @@ -20,12 +18,11 @@ pub(crate) fn azure_auth_inputs(request: &PreparedOcrRequest) -> Result Option + Sync), ) -> Result>, Error> { - static SERVICE: OnceLock = OnceLock::new(); - SERVICE - .get_or_init(AzureAuthService::default) + service .get_azure_ad_token(config, env_lookup) .await .or_else(|error| match error { diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs index 51a2668310e..5ee5ab3be94 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs @@ -183,12 +183,15 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOcrConfig { async fn validate_environment( &self, request: &PreparedOcrRequest, - _client: &OcrClient, + client: &OcrClient, ) -> Result { let config = crate::azure_ai::ocr::common_utils::azure_auth_inputs(request)?; - self.resolve_headers(&request.connection, &config, &|name: &str| { - request.connection.secret(name) - }) + self.resolve_headers( + &client.auth().azure, + &request.connection, + &config, + &|name: &str| request.connection.secret(name), + ) .await } @@ -450,7 +453,7 @@ fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result Option + Sync), @@ -635,7 +639,7 @@ impl AzureDocumentIntelligenceOcrConfig { .collect(), ); } - let token = super::super::common_utils::resolve_entra(config, env_lookup) + let token = super::super::common_utils::resolve_entra(auth, config, env_lookup) .await? .ok_or(Error::MissingAzureDocumentIntelligenceCredentials)?; super::super::common_utils::validate_destination(connection, token.source())?; @@ -809,9 +813,12 @@ mod tests { }; let error = AzureDocumentIntelligenceOcrConfig - .resolve_headers(&connection, &Default::default(), &|name| { - (name == AZURE_DI_API_KEY_ENV).then(|| "environment-key".into()) - }) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|name| (name == AZURE_DI_API_KEY_ENV).then(|| "environment-key".into()), + ) .await .unwrap_err(); @@ -833,7 +840,12 @@ mod tests { }; let headers = AzureDocumentIntelligenceOcrConfig - .resolve_headers(&connection, &Default::default(), &|_| None) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|_| None, + ) .await .unwrap(); diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs index 1556ae2a414..8efc27b0bea 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs @@ -60,12 +60,15 @@ impl BaseOcrConfig for AzureAiOcrConfig { async fn validate_environment( &self, request: &PreparedOcrRequest, - _client: &OcrClient, + client: &OcrClient, ) -> Result { let config = crate::azure_ai::ocr::common_utils::azure_auth_inputs(request)?; - self.resolve_headers(&request.connection, &config, &|name: &str| { - request.connection.secret(name) - }) + self.resolve_headers( + &client.auth().azure, + &request.connection, + &config, + &|name: &str| request.connection.secret(name), + ) .await } @@ -139,6 +142,7 @@ impl AzureAiOcrConfig { async fn resolve_headers( &self, + auth: &litellm_auth_azure::AzureAuthService, connection: &OcrConnection, config: &AzureAuthInputs, env_lookup: &(dyn Fn(&str) -> Option + Sync), @@ -146,7 +150,7 @@ impl AzureAiOcrConfig { Self::resolve_api_base(connection.api_base.as_deref(), env_lookup)?; if litellm_http::request::has_header(&connection.extra_headers, "authorization") { if config.azure_ad_token_provider.is_some() { - super::common_utils::resolve_entra(config, env_lookup).await?; + super::common_utils::resolve_entra(auth, config, env_lookup).await?; } super::common_utils::validate_destination(connection, connection.extra_headers_source)?; return Ok(connection.extra_headers.clone()); @@ -166,7 +170,7 @@ impl AzureAiOcrConfig { super::common_utils::validate_destination(connection, key.source())?; return Ok(bearer_headers(connection, key.value())); } - let key = super::common_utils::resolve_entra(config, env_lookup) + let key = super::common_utils::resolve_entra(auth, config, env_lookup) .await? .ok_or(Error::MissingAzureAiCredentials)?; super::common_utils::validate_destination(connection, key.source())?; @@ -253,9 +257,12 @@ mod tests { }; assert_eq!( AzureAiOcrConfig - .resolve_headers(&connection, &Default::default(), &|_| { - Some("environment-key".into()) - }) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|_| { Some("environment-key".into()) } + ) .await .unwrap(), connection.extra_headers @@ -267,9 +274,12 @@ mod tests { async fn request_key_precedes_environment_key(connection: OcrConnection) { assert_eq!( AzureAiOcrConfig - .resolve_headers(&connection, &Default::default(), &|_| { - Some("environment-key".into()) - }) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|_| { Some("environment-key".into()) } + ) .await .unwrap()[0], ("Authorization".into(), "Bearer request-key".into()) @@ -285,9 +295,12 @@ mod tests { }; let error = AzureAiOcrConfig - .resolve_headers(&connection, &Default::default(), &|name| { - (name == AZURE_AI_API_KEY_ENV).then(|| "environment-key".into()) - }) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|name| (name == AZURE_AI_API_KEY_ENV).then(|| "environment-key".into()), + ) .await .unwrap_err(); @@ -309,7 +322,12 @@ mod tests { }; let headers = AzureAiOcrConfig - .resolve_headers(&connection, &Default::default(), &|_| None) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|_| None, + ) .await .unwrap(); @@ -329,7 +347,7 @@ mod tests { let connection = OcrConnection::default(); let headers = AzureAiOcrConfig - .resolve_headers(&connection, &Default::default(), &env) + .resolve_headers(&Default::default(), &connection, &Default::default(), &env) .await .unwrap(); let url = AzureAiOcrConfig.build_ocr_url(None, &env).unwrap(); diff --git a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/mod.rs b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/mod.rs index f239b6921fa..fa7df180f50 100644 --- a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/mod.rs +++ b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/mod.rs @@ -1 +1,2 @@ +pub mod streaming; pub mod transformation; diff --git a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/streaming.rs b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/streaming.rs new file mode 100644 index 00000000000..abb61297669 --- /dev/null +++ b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/streaming.rs @@ -0,0 +1,141 @@ +use bytes::Bytes; +use futures_util::{StreamExt, stream::BoxStream}; +use litellm_framing::{frames, sse::SseCodec}; + +pub use crate::base_llm::base_model_iterator::ByteStream; +use crate::{Error, anthropic::messages::streaming_iterator::AnthropicMessagesStreamEvent}; + +pub type EventStream = BoxStream<'static, Result>; +pub type StreamDecoder = fn(ByteStream) -> EventStream; + +pub fn anthropic_sse_event_stream(bytes: ByteStream) -> EventStream { + Box::pin(frames(bytes, SseCodec::default()).map(|event| { + let event = event + .map_err(|error| Error::InvalidResponse(format!("stream framing failed: {error}")))?; + serde_json::from_str(&event.data).map_err(|error| { + Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}")) + }) + })) +} + +pub fn encode_anthropic_sse(event: &AnthropicMessagesStreamEvent) -> Result { + let data = serde_json::to_value(event).map_err(|error| { + Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}")) + })?; + let name = data + .get("type") + .and_then(serde_json::Value::as_str) + .ok_or_else(|| { + Error::InvalidResponse( + "Anthropic stream event is invalid: stream event has no type".into(), + ) + })?; + Ok(Bytes::from(format!("event: {name}\ndata: {data}\n\n"))) +} + +#[cfg(test)] +mod tests { + use futures_util::{StreamExt, TryStreamExt, stream}; + use serde_json::json; + + use super::*; + use crate::anthropic::messages::streaming_iterator::{ + AnthropicContentBlockDelta, AnthropicStreamUsage, + }; + + const TEXT_DELTA: &str = + r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hello"}}"#; + + fn in_pieces(wire: &[u8]) -> ByteStream { + let pieces: Vec = wire.chunks(3).map(Bytes::copy_from_slice).collect(); + stream::iter(pieces.into_iter().map(Ok)).boxed() + } + + #[tokio::test] + async fn sse_frames_split_anywhere_decode_into_typed_events() { + let wire = format!("event: content_block_delta\ndata: {TEXT_DELTA}\n\n"); + let events = anthropic_sse_event_stream(in_pieces(wire.as_bytes())) + .try_collect::>() + .await + .unwrap(); + + assert_eq!( + events, + vec![AnthropicMessagesStreamEvent::ContentBlockDelta { + index: 0, + delta: AnthropicContentBlockDelta::TextDelta { + text: "hello".into(), + }, + }] + ); + } + + #[tokio::test] + async fn decodes_citations_delta_events() { + let wire = concat!( + "event: content_block_delta\n", + r#"data: {"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"#, + "\n\n", + ); + let events = anthropic_sse_event_stream(in_pieces(wire.as_bytes())) + .try_collect::>() + .await + .unwrap(); + + assert!(matches!( + events.as_slice(), + [AnthropicMessagesStreamEvent::ContentBlockDelta { + delta: AnthropicContentBlockDelta::Citations { .. }, + .. + }] + )); + } + + fn events() -> Vec { + vec![ + AnthropicMessagesStreamEvent::Ping, + AnthropicMessagesStreamEvent::ContentBlockDelta { + index: 1, + delta: AnthropicContentBlockDelta::TextDelta { text: "hi".into() }, + }, + AnthropicMessagesStreamEvent::ContentBlockStop { index: 1 }, + AnthropicMessagesStreamEvent::MessageStop { + usage: Some(AnthropicStreamUsage { + output_tokens: Some(7), + ..AnthropicStreamUsage::default() + }), + }, + ] + } + + #[tokio::test] + async fn encoded_events_decode_back_to_themselves() { + let wire = events() + .iter() + .map(encode_anthropic_sse) + .collect::, _>>() + .unwrap(); + + let decoded = anthropic_sse_event_stream(stream::iter(wire.into_iter().map(Ok)).boxed()) + .try_collect::>() + .await + .unwrap(); + + assert_eq!(decoded, events()); + } + + #[test] + fn an_event_is_named_by_its_type() { + let encoded = + encode_anthropic_sse(&AnthropicMessagesStreamEvent::MessageStop { usage: None }) + .unwrap(); + + assert_eq!( + encoded, + Bytes::from(format!( + "event: message_stop\ndata: {}\n\n", + json!({"type": "message_stop"}) + )) + ); + } +} diff --git a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs index 8db14687214..9e9f585263c 100644 --- a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs @@ -1,30 +1,13 @@ -use litellm_http::request::{has_bearer_auth, has_header}; use litellm_types::llms::anthropic_messages::{ anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse, }; +pub use crate::base_llm::auth::{Headers, ValidatedEnvironment}; use crate::{ - anthropic::experimental_pass_through::messages::thinking::ThinkingContext, - base_llm::chat::transformation::Error, + Error, anthropic::messages::thinking::ThinkingContext, + base_llm::anthropic_messages::streaming::StreamDecoder, }; -pub type Headers = Vec<(String, String)>; - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum MessagesAuthStrategy { - Bearer, - Header(&'static str), -} - -impl MessagesAuthStrategy { - pub fn header_name(self) -> &'static str { - match self { - Self::Bearer => "authorization", - Self::Header(header_name) => header_name, - } - } -} - #[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] pub struct MessagesTransformContext { pub thinking: ThinkingContext, @@ -39,6 +22,15 @@ pub trait BaseAnthropicMessagesConfig: Sync { env_lookup: &dyn Fn(&str) -> Option, ) -> Result; + fn complete_stream_url( + &self, + api_base: Option<&str>, + model: &str, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + self.get_complete_url(api_base, model, env_lookup) + } + fn transform_anthropic_messages_request( &self, request: AnthropicMessagesRequest, @@ -55,42 +47,24 @@ pub trait BaseAnthropicMessagesConfig: Sync { Ok(response) } - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; - fn secret_names(&self) -> &'static [&'static str]; - fn auth_strategy(&self) -> MessagesAuthStrategy { - MessagesAuthStrategy::Header("x-api-key") - } - - fn accepts_bearer_auth(&self) -> bool { - false - } - - fn authenticate( + /// Shapes the forwarded headers and names the credential, the way Python's + /// `validate_environment` does, without applying it: `resolve_auth` does that once + /// for every config. + fn validate_environment( &self, headers: Headers, api_key: Option<&str>, + model: &str, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - let strategy = self.auth_strategy(); - if has_header(&headers, strategy.header_name()) - || (self.accepts_bearer_auth() && has_bearer_auth(&headers)) - { - return Ok(headers); - } - let api_key = self.resolve_api_key(api_key, env_lookup)?; - let auth_header = match strategy { - MessagesAuthStrategy::Bearer => { - ("authorization".to_string(), format!("Bearer {api_key}")) - } - MessagesAuthStrategy::Header(name) => (name.to_string(), api_key), - }; - Ok(headers.into_iter().chain([auth_header]).collect()) + ) -> Result; + + /// `None` relays the upstream bytes untouched, which is right for every host that already + /// speaks Anthropic SSE. A host on another wire returns the decoder that lifts its frames + /// into Anthropic stream events, and the route re-encodes those as Anthropic SSE. + fn stream_decoder(&self) -> Option { + None } fn default_headers(&self) -> &'static [(&'static str, &'static str)] { @@ -107,49 +81,8 @@ pub trait BaseAnthropicMessagesConfig: Sync { #[cfg(test)] mod tests { - use rstest::rstest; - use super::*; - - const X_API_KEY: MessagesAuthStrategy = MessagesAuthStrategy::Header("x-api-key"); - - struct StubConfig { - strategy: MessagesAuthStrategy, - accepts_bearer: bool, - } - - impl BaseAnthropicMessagesConfig for StubConfig { - fn secret_names(&self) -> &'static [&'static str] { - &[] - } - - fn get_complete_url( - &self, - _api_base: Option<&str>, - _model: &str, - _env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - Ok(String::new()) - } - - fn resolve_api_key( - &self, - api_key: Option<&str>, - _env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - api_key - .map(str::to_string) - .ok_or(Error::MissingField("api_key")) - } - - fn auth_strategy(&self) -> MessagesAuthStrategy { - self.strategy - } - - fn accepts_bearer_auth(&self) -> bool { - self.accepts_bearer - } - } + use crate::base_llm::auth::AuthScheme; struct DefaultsConfig; @@ -167,32 +100,20 @@ mod tests { Ok(String::new()) } - fn resolve_api_key( + fn validate_environment( &self, - api_key: Option<&str>, + headers: Headers, + _api_key: Option<&str>, + _model: &str, _env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - api_key - .map(str::to_string) - .ok_or(Error::MissingField("api_key")) + ) -> Result { + Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Forwarded, + }) } } - #[test] - fn default_config_adds_its_key_next_to_a_forwarded_bearer() { - assert_eq!( - DefaultsConfig.authenticate( - headers(&[("authorization", "Bearer forwarded")]), - Some("sk"), - &|_| None - ), - Ok(headers(&[ - ("authorization", "Bearer forwarded"), - ("x-api-key", "sk") - ])) - ); - } - #[test] fn default_request_headers_are_the_given_headers() { let request: AnthropicMessagesRequest = serde_json::from_value(serde_json::json!({ @@ -214,82 +135,4 @@ mod tests { .map(|(name, value)| (name.to_string(), value.to_string())) .collect() } - - #[rstest] - #[case::own_header_is_kept( - X_API_KEY, - false, - headers(&[("x-api-key", "forwarded")]), - None, - Ok(headers(&[("x-api-key", "forwarded")])) - )] - #[case::own_header_in_any_casing_is_kept( - X_API_KEY, - false, - headers(&[("X-Api-Key", "forwarded")]), - None, - Ok(headers(&[("X-Api-Key", "forwarded")])) - )] - #[case::accepted_bearer_is_kept( - X_API_KEY, - true, - headers(&[("authorization", "Bearer forwarded")]), - None, - Ok(headers(&[("authorization", "Bearer forwarded")])) - )] - #[case::bearer_the_provider_does_not_accept_gets_the_key_too( - X_API_KEY, - false, - headers(&[("authorization", "Bearer forwarded")]), - Some("sk"), - Ok(headers(&[("authorization", "Bearer forwarded"), ("x-api-key", "sk")])) - )] - #[case::blank_bearer_gets_the_key( - X_API_KEY, - true, - headers(&[("authorization", "Bearer ")]), - Some("sk"), - Ok(headers(&[("authorization", "Bearer "), ("x-api-key", "sk")])) - )] - #[case::key_goes_in_the_provider_header( - X_API_KEY, - false, - headers(&[("content-type", "application/json")]), - Some("sk"), - Ok(headers(&[("content-type", "application/json"), ("x-api-key", "sk")])) - )] - #[case::key_goes_in_a_bearer( - MessagesAuthStrategy::Bearer, - false, - headers(&[]), - Some("sk"), - Ok(headers(&[("authorization", "Bearer sk")])) - )] - #[case::bearer_strategy_keeps_a_forwarded_authorization( - MessagesAuthStrategy::Bearer, - false, - headers(&[("authorization", "Bearer forwarded")]), - None, - Ok(headers(&[("authorization", "Bearer forwarded")])) - )] - #[case::missing_key_is_an_error( - X_API_KEY, - false, - headers(&[]), - None, - Err(Error::MissingField("api_key")) - )] - fn default_authenticate_applies_the_key_unless_a_credential_is_forwarded( - #[case] strategy: MessagesAuthStrategy, - #[case] accepts_bearer: bool, - #[case] forwarded: Headers, - #[case] api_key: Option<&str>, - #[case] expected: Result, - ) { - let config = StubConfig { - strategy, - accepts_bearer, - }; - assert_eq!(config.authenticate(forwarded, api_key, &|_| None), expected); - } } diff --git a/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs b/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs index 1257bbf0d6a..562902ac6a8 100644 --- a/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs @@ -1,7 +1,7 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; -use crate::base_llm::chat::transformation::Error; +use crate::Error; #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct AudioTranscriptionRequestData { @@ -21,7 +21,7 @@ impl AudioTranscriptionResponseData { } } -pub use litellm_auth::RequestAuth; +pub use crate::base_llm::auth::{Headers, ValidatedEnvironment}; pub trait BaseAudioTranscriptionConfig: Sync { fn get_supported_openai_params(&self) -> &'static [&'static str]; @@ -58,10 +58,11 @@ pub trait BaseAudioTranscriptionConfig: Sync { response_json: Value, ) -> Result; - fn auth_strategy( + fn validate_environment( &self, + headers: Headers, model: &str, optional_params: &Map, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; + ) -> Result; } diff --git a/litellm-rust/crates/llms/src/base_llm/auth.rs b/litellm-rust/crates/llms/src/base_llm/auth.rs new file mode 100644 index 00000000000..7d1b0014dcd --- /dev/null +++ b/litellm-rust/crates/llms/src/base_llm/auth.rs @@ -0,0 +1,263 @@ +//! How a provider call authenticates, decided by the provider config when the request is +//! prepared and applied once here when it is sent. +//! +//! Python folds this into `validate_environment` plus `sign_request`. The Rust configs keep +//! that split: `validate_environment` shapes the forwarded headers and names the credential +//! as an [`AuthScheme`], and [`resolve_auth`] turns the scheme into headers and a signer. + +use litellm_auth::{AuthServices, CredentialPlacement, SecretValue, TokenProviderHandle}; +use litellm_auth_aws::{AwsCredentialSource, SigV4Signer}; +use litellm_http::request::without_headers; + +pub type Headers = Vec<(String, String)>; + +#[derive(Clone, Debug)] +pub enum AuthScheme { + /// The caller's own credential is already in the headers and is sent as is. + Forwarded, + /// A credential in hand, placed in its header. A forwarded header of the same name is + /// replaced: the deployment's identity outranks the caller's. + Credential { + placement: CredentialPlacement, + secret: SecretValue, + }, + /// A bearer acquired when the request is sent, from a token source such as a cloud SDK + /// or a caller-supplied callable. + Token { provider: TokenProviderHandle }, + /// AWS SigV4 over the bytes that go on the wire, so the handler signs after the body is + /// serialized. + AwsSigV4 { + region: String, + service: &'static str, + credentials: Box, + }, +} + +/// The outcome of a config's `validate_environment`: the headers it shaped and how the +/// call authenticates. +#[derive(Clone, Debug)] +pub struct ValidatedEnvironment { + pub headers: Headers, + pub auth: AuthScheme, +} + +#[derive(Debug)] +pub struct Authenticated { + pub headers: Headers, + pub signer: Option, +} + +pub async fn resolve_auth( + services: &AuthServices, + validated: ValidatedEnvironment, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Result { + let ValidatedEnvironment { headers, auth } = validated; + match auth { + AuthScheme::Forwarded => Ok(Authenticated { + headers, + signer: None, + }), + AuthScheme::Credential { placement, secret } => Ok(Authenticated { + headers: with_credential(headers, placement, secret.expose()), + signer: None, + }), + AuthScheme::Token { provider } => { + let token = provider.acquire().await?; + Ok(Authenticated { + headers: with_credential( + headers, + CredentialPlacement::Bearer, + token.secret().expose(), + ), + signer: None, + }) + } + AuthScheme::AwsSigV4 { + region, + service, + credentials, + } => Ok(Authenticated { + headers, + signer: Some( + SigV4Signer::resolve(&services.aws, region, service, *credentials, env_lookup) + .await?, + ), + }), + } +} + +/// Fills in the defaults the caller did not forward, matching Python's +/// `if name not in headers` checks. +pub fn with_default_headers(headers: Headers, defaults: &[(&str, &str)]) -> Headers { + let missing: Vec<(String, String)> = defaults + .iter() + .filter(|(name, _)| { + !headers + .iter() + .any(|(header, _)| header.eq_ignore_ascii_case(name)) + }) + .map(|(name, value)| ((*name).to_string(), (*value).to_string())) + .collect(); + headers.into_iter().chain(missing).collect() +} + +fn with_credential(headers: Headers, placement: CredentialPlacement, credential: &str) -> Headers { + let name = placement.header_name(); + let value = match placement { + CredentialPlacement::Bearer => format!("Bearer {credential}"), + CredentialPlacement::Header(_) => credential.to_string(), + }; + without_headers(headers, &[name]) + .into_iter() + .chain([(name.to_ascii_lowercase(), value)]) + .collect() +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use litellm_auth::{AuthServices, ResolvedCredential, TokenFuture, TokenProvider}; + use litellm_auth_aws::Credentials; + use rstest::rstest; + + use super::*; + + fn no_env(_: &str) -> Option { + None + } + + fn headers(pairs: &[(&str, &str)]) -> Headers { + pairs + .iter() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect() + } + + async fn resolve(headers: Headers, auth: AuthScheme) -> Authenticated { + resolve_auth( + &AuthServices::default(), + ValidatedEnvironment { headers, auth }, + &no_env, + ) + .await + .unwrap() + } + + #[rstest] + #[case::header_is_appended( + &[("content-type", "application/json")], + CredentialPlacement::Header("x-api-key"), + &[("content-type", "application/json"), ("x-api-key", "sk")], + )] + #[case::forwarded_header_of_the_same_name_is_replaced_in_any_casing( + &[("X-Api-Key", "caller"), ("x-trace", "1")], + CredentialPlacement::Header("x-api-key"), + &[("x-trace", "1"), ("x-api-key", "sk")], + )] + #[case::bearer_replaces_a_forwarded_authorization( + &[("Authorization", "Bearer caller")], + CredentialPlacement::Bearer, + &[("authorization", "Bearer sk")], + )] + #[tokio::test] + async fn a_credential_lands_in_its_header_and_outranks_the_forwarded_one( + #[case] forwarded: &[(&str, &str)], + #[case] placement: CredentialPlacement, + #[case] expected: &[(&str, &str)], + ) { + let authenticated = resolve( + headers(forwarded), + AuthScheme::Credential { + placement, + secret: SecretValue::new("sk"), + }, + ) + .await; + assert_eq!(authenticated.headers, headers(expected)); + assert!(authenticated.signer.is_none()); + } + + #[rstest] + #[case::nothing_forwarded( + &[], + &[("x-version", "1"), ("content-type", "application/json")], + &[("x-version", "1"), ("content-type", "application/json")], + )] + #[case::forwarded_header_wins_in_any_case( + &[("X-Version", "custom"), ("x-api-key", "k")], + &[("x-version", "1"), ("content-type", "application/json")], + &[("X-Version", "custom"), ("x-api-key", "k"), ("content-type", "application/json")], + )] + #[case::no_defaults(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])] + fn default_headers_fill_only_missing_names( + #[case] forwarded: &[(&str, &str)], + #[case] defaults: &[(&str, &str)], + #[case] expected: &[(&str, &str)], + ) { + assert_eq!( + with_default_headers(headers(forwarded), defaults), + headers(expected) + ); + } + + #[tokio::test] + async fn forwarded_auth_sends_the_headers_untouched() { + let forwarded = headers(&[("x-api-key", "caller"), ("authorization", "Bearer caller")]); + let authenticated = resolve(forwarded.clone(), AuthScheme::Forwarded).await; + assert_eq!(authenticated.headers, forwarded); + assert!(authenticated.signer.is_none()); + } + + #[derive(Debug)] + struct StaticToken(&'static str); + + impl TokenProvider for StaticToken { + fn acquire(&self) -> TokenFuture<'_> { + Box::pin(async move { + Ok(ResolvedCredential::AccessToken { + token: SecretValue::new(self.0), + expires_on: None, + }) + }) + } + } + + #[tokio::test] + async fn a_token_is_acquired_at_send_time_and_sent_as_a_bearer() { + let authenticated = resolve( + headers(&[("authorization", "Bearer stale")]), + AuthScheme::Token { + provider: TokenProviderHandle::new(Arc::new(StaticToken("fresh"))), + }, + ) + .await; + assert_eq!( + authenticated.headers, + headers(&[("authorization", "Bearer fresh")]) + ); + } + + #[tokio::test] + async fn sigv4_leaves_the_headers_to_the_signer() { + let forwarded = headers(&[("x-request-id", "abc")]); + let authenticated = resolve( + forwarded.clone(), + AuthScheme::AwsSigV4 { + region: "us-east-1".into(), + service: "bedrock", + credentials: Box::new(AwsCredentialSource::HostSupplied(Credentials::new( + "AKIDEXAMPLE", + "secret", + None, + None, + "test", + ))), + }, + ) + .await; + assert_eq!(authenticated.headers, forwarded); + assert!(authenticated.signer.is_some()); + } +} diff --git a/litellm-rust/crates/llms/src/base_llm/base_model_iterator.rs b/litellm-rust/crates/llms/src/base_llm/base_model_iterator.rs index 928ef80b29a..a283ca84089 100644 --- a/litellm-rust/crates/llms/src/base_llm/base_model_iterator.rs +++ b/litellm-rust/crates/llms/src/base_llm/base_model_iterator.rs @@ -1,3 +1,10 @@ +use std::{collections::VecDeque, convert::Infallible, io, pin::Pin}; + +use bytes::Bytes; +use futures_util::{Stream, StreamExt, stream, stream::BoxStream}; + +pub type ByteStream = BoxStream<'static, Result>; + pub trait StreamTransformer { type Input; type Output; @@ -7,3 +14,149 @@ pub trait StreamTransformer { fn finish(&mut self) -> Result, Self::Error>; } + +#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] +pub enum StreamError { + #[error(transparent)] + Decode(D), + #[error(transparent)] + Transform(T), +} + +impl StreamError { + pub fn into_decode(self) -> D { + match self { + Self::Decode(error) => error, + Self::Transform(never) => match never {}, + } + } +} + +struct Driver { + events: Pin>, + transformer: T, + ready: VecDeque, + finished: bool, +} + +/// Drives `transformer` over `events`, then flushes it with `finish`. The first error ends the +/// stream. +pub fn transform_stream( + events: S, + transformer: T, +) -> impl Stream>> + Send +where + S: Stream> + Send, + T: StreamTransformer + Send, + T::Output: Send, + T::Error: Send, + D: Send, +{ + let driver = Driver { + events: Box::pin(events), + transformer, + ready: VecDeque::new(), + finished: false, + }; + stream::unfold(driver, |mut driver| async move { + loop { + if let Some(output) = driver.ready.pop_front() { + return Some((Ok(output), driver)); + } + if driver.finished { + return None; + } + match driver.events.next().await { + Some(Ok(event)) => match driver.transformer.transform(event) { + Ok(outputs) => driver.ready.extend(outputs), + Err(error) => { + driver.finished = true; + return Some((Err(StreamError::Transform(error)), driver)); + } + }, + Some(Err(error)) => { + driver.finished = true; + return Some((Err(StreamError::Decode(error)), driver)); + } + None => { + driver.finished = true; + match driver.transformer.finish() { + Ok(outputs) => driver.ready.extend(outputs), + Err(error) => return Some((Err(StreamError::Transform(error)), driver)), + } + } + } + } + }) +} + +#[cfg(test)] +mod tests { + use futures_util::TryStreamExt; + + use super::*; + + struct Doubler; + + impl StreamTransformer for Doubler { + type Input = u32; + type Output = u32; + type Error = String; + + fn transform(&mut self, input: u32) -> Result, String> { + match input { + 0 => Err("zero".into()), + n => Ok(vec![n, n * 2]), + } + } + + fn finish(&mut self) -> Result, String> { + Ok(vec![u32::MAX]) + } + } + + #[tokio::test] + async fn flat_maps_each_event_and_flushes_at_the_end() { + let output = transform_stream(stream::iter([Ok::<_, String>(1), Ok(2)]), Doubler) + .try_collect::>() + .await + .unwrap(); + + assert_eq!(output, vec![1, 2, 2, 4, u32::MAX]); + } + + #[tokio::test] + async fn a_transform_error_ends_the_stream_without_flushing() { + let output = transform_stream(stream::iter([Ok::<_, String>(1), Ok(0), Ok(3)]), Doubler) + .collect::>() + .await; + + assert_eq!( + output, + vec![ + Ok(1), + Ok(2), + Err(StreamError::Transform("zero".to_string())) + ] + ); + } + + #[tokio::test] + async fn a_decode_error_ends_the_stream_without_flushing() { + let output = transform_stream( + stream::iter([Ok(1), Err("bad frame".to_string()), Ok(3)]), + Doubler, + ) + .collect::>() + .await; + + assert_eq!( + output, + vec![ + Ok(1), + Ok(2), + Err(StreamError::Decode("bad frame".to_string())) + ] + ); + } +} diff --git a/litellm-rust/crates/llms/src/base_llm/chat/mod.rs b/litellm-rust/crates/llms/src/base_llm/chat/mod.rs index f239b6921fa..fa7df180f50 100644 --- a/litellm-rust/crates/llms/src/base_llm/chat/mod.rs +++ b/litellm-rust/crates/llms/src/base_llm/chat/mod.rs @@ -1 +1,2 @@ +pub mod streaming; pub mod transformation; diff --git a/litellm-rust/crates/llms/src/base_llm/chat/streaming.rs b/litellm-rust/crates/llms/src/base_llm/chat/streaming.rs new file mode 100644 index 00000000000..b9d715bcd68 --- /dev/null +++ b/litellm-rust/crates/llms/src/base_llm/chat/streaming.rs @@ -0,0 +1,54 @@ +use std::collections::HashMap; + +use futures_util::{StreamExt, stream::BoxStream}; +use litellm_types::utils::ChatCompletionChunk; + +use crate::{ + Error, + base_llm::base_model_iterator::{ByteStream, StreamError, StreamTransformer, transform_stream}, +}; + +pub type ChatChunkStream = BoxStream<'static, Result>; + +/// What Python's `map_openai_params` decides about the stream and `completion` +/// hands to `ModelResponseIterator`: it is settled while the request is built, +/// never re-derived from the body. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct StreamShape { + pub json_mode: bool, + pub speed: Option, + pub tool_name_reverse_map: HashMap, +} + +/// A wire decoder paired with the iterator that turns its events into chat chunks. +/// A config names both; the core runs the pair over the response bytes. +pub struct ChatStream { + run: Box ChatChunkStream + Send>, +} + +impl ChatStream { + pub fn new( + decode: fn(ByteStream) -> BoxStream<'static, Result>, + iterator: T, + ) -> Self + where + E: Send + 'static, + T: StreamTransformer + + Send + + 'static, + { + Self { + run: Box::new(move |bytes| { + Box::pin(transform_stream(decode(bytes), iterator).map(|item| { + item.map_err(|error| match error { + StreamError::Decode(error) | StreamError::Transform(error) => error, + }) + })) + }), + } + } + + pub fn run(self, bytes: ByteStream) -> ChatChunkStream { + (self.run)(bytes) + } +} diff --git a/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs b/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs index c7d1a27c71e..8a074e59207 100644 --- a/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs @@ -4,30 +4,17 @@ use litellm_types::{ }; use serde_json::{Map, Value}; -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("expected {expected}, got {actual}")] - InvalidType { - expected: &'static str, - actual: &'static str, - }, - #[error("missing required field: {0}")] - MissingField(&'static str), - #[error("invalid request: {0}")] - InvalidRequest(String), - #[error("invalid response: {0}")] - InvalidResponse(String), - #[error("unsupported: {0}")] - Unsupported(&'static str), - #[error(transparent)] - Auth(#[from] litellm_auth::Error), -} +use crate::{ + Error, + base_llm::chat::streaming::{ChatStream, StreamShape}, +}; /// The provider-shaped request body a config produces. Named rather than a bare /// `Value` so the transform contract stays a typed one, mirroring /// [`crate::base_llm::audio_transcription::transformation::AudioTranscriptionRequestData`]. pub struct ProviderChatRequestData { pub body: Value, + pub stream_shape: StreamShape, } /// The raw provider response body handed back to a config for normalization. @@ -41,7 +28,7 @@ pub const STREAM_PARAM: &str = "stream"; /// presence does not make a request untranslatable. const IGNORABLE_MESSAGE_FIELDS: &[&str] = &["name"]; -pub use litellm_auth::RequestAuth; +pub use crate::base_llm::auth::{Headers, ValidatedEnvironment}; /// Why a request cannot be served by the Rust path. /// @@ -78,28 +65,27 @@ pub trait BaseConfig: Sync { response: ProviderChatResponseData, ) -> Result; - fn auth( + /// `None` means this config has no streaming path yet, so the host keeps the request. + fn model_response_iterator(&self, _shape: StreamShape) -> Option { + None + } + + /// Shapes the forwarded headers and names the credential, the way Python's + /// `validate_environment` does, without applying it: `resolve_auth` does that once + /// for every config. + fn validate_environment( &self, + headers: Headers, api_key: Option<&str>, model: &str, optional_params: &Map, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; + ) -> Result; fn default_headers(&self) -> &'static [(&'static str, &'static str)] { &[("content-type", "application/json")] } - /// Whether an auth header the caller already supplied is the credential this - /// request should authenticate with, so the resolved one is not applied. - /// - /// Defaults to false: the deployment's credential outranks anything - /// forwarded, which is what every provider wants for its own auth header. - /// A provider overrides this only for a scheme it hands off to entirely. - fn defers_to_forwarded_auth(&self, _headers: &[(String, String)]) -> bool { - false - } - /// Parameters consumed as call configuration (credentials, endpoints) /// rather than placed in the body. Accepted, never serialized. fn config_params(&self) -> &'static [&'static str] { diff --git a/litellm-rust/crates/llms/src/base_llm/mod.rs b/litellm-rust/crates/llms/src/base_llm/mod.rs index 8ed37da4573..399b932e9da 100644 --- a/litellm-rust/crates/llms/src/base_llm/mod.rs +++ b/litellm-rust/crates/llms/src/base_llm/mod.rs @@ -1,5 +1,6 @@ pub mod anthropic_messages; pub mod audio_transcription; +pub mod auth; pub mod base_model_iterator; pub mod chat; pub mod ocr; diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/document.rs b/litellm-rust/crates/llms/src/base_llm/ocr/document.rs index 724625b8208..9bcaad353ab 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/document.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/document.rs @@ -184,16 +184,21 @@ mod tests { reqwest::header::AUTHORIZATION, reqwest::header::HeaderValue::from_static("Bearer provider-secret"), ); - let provider_http = reqwest::Client::builder() - .default_headers(provider_headers) - .build() - .unwrap(); - let document_http = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .unwrap(); - let client = - crate::base_llm::ocr::handler::OcrClient::for_test(provider_http, document_http); + #[expect( + clippy::disallowed_methods, + clippy::disallowed_types, + reason = "the pool has no default-header setting to stand in for provider credentials" + )] + let provider_http = litellm_http::Client::for_test( + reqwest::Client::builder() + .default_headers(provider_headers) + .build() + .unwrap(), + ); + let client = crate::base_llm::ocr::handler::OcrClient::for_test( + provider_http, + litellm_http::Client::no_redirect_for_test(), + ); let converted = inline_remote_document( client.document_fetcher(), OcrDocument::ImageUrl { diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs index 9527fd20f2d..0148ca2841b 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs @@ -2,10 +2,10 @@ use std::sync::Arc; use bytes::{Bytes, BytesMut}; use futures_util::future::BoxFuture; -use litellm_auth_gcp::VertexAuth; +use litellm_auth::AuthServices; use litellm_host::event::WireRequest; use litellm_http::{ - ClientVariant, HttpClientConfig, HttpClientPool, + Client, ClientVariant, HttpClientConfig, HttpClientPool, media::{MediaFetcher, UrlPolicy}, outbound::{OutboundRequest, RequestSigner}, transport, @@ -33,10 +33,10 @@ pub trait CallHooks: Send + Sync { #[derive(Clone)] pub struct OcrClient { - provider_http: reqwest::Client, - polling_http: reqwest::Client, + provider_http: Client, + polling_http: Client, document_fetcher: MediaFetcher, - vertex_auth: VertexAuth, + auth: Arc, settings: OcrSettings, secrets: Arc, } @@ -46,7 +46,7 @@ impl OcrClient { pool: &HttpClientPool, config: &HttpClientConfig, url_policy: UrlPolicy, - vertex_auth: VertexAuth, + auth: Arc, settings: OcrSettings, secrets: Arc, ) -> Result { @@ -54,17 +54,17 @@ impl OcrClient { provider_http: pool.client(config, ClientVariant::Provider)?, polling_http: pool.client(config, ClientVariant::NoRedirect)?, document_fetcher: MediaFetcher::new(pool, config, url_policy)?, - vertex_auth, + auth, settings, secrets, }) } - pub fn provider_http(&self) -> &reqwest::Client { + pub fn provider_http(&self) -> &Client { &self.provider_http } - pub fn polling_http(&self) -> &reqwest::Client { + pub fn polling_http(&self) -> &Client { &self.polling_http } @@ -72,8 +72,8 @@ impl OcrClient { &self.document_fetcher } - pub fn vertex_auth(&self) -> &VertexAuth { - &self.vertex_auth + pub fn auth(&self) -> &AuthServices { + &self.auth } pub fn settings(&self) -> &OcrSettings { @@ -85,17 +85,18 @@ impl OcrClient { } #[cfg(any(test, feature = "test-support"))] - pub fn for_test(provider_http: reqwest::Client, document_http: reqwest::Client) -> Self { + pub fn for_test(provider_http: Client, no_redirect_http: Client) -> Self { Self { + secrets: Arc::new( + litellm_secrets::source::EnvironmentSecrets::python_compatible( + provider_http.clone(), + ), + ), provider_http, - polling_http: reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("test polling client builds"), - document_fetcher: MediaFetcher::for_test(document_http), - vertex_auth: VertexAuth::default(), + polling_http: no_redirect_http.clone(), + document_fetcher: MediaFetcher::for_test(no_redirect_http), + auth: Arc::new(AuthServices::default()), settings: OcrSettings::default(), - secrets: Arc::new(litellm_secrets::source::EnvironmentSecrets::default()), } } @@ -311,7 +312,7 @@ mod tests { let _connection = listener.accept().await.unwrap(); tokio::time::sleep(Duration::from_secs(1)).await; }); - let error = reqwest::Client::new() + let error = litellm_http::Client::plain_for_test() .get(format!("http://{address}")) .timeout(Duration::from_millis(10)) .send() diff --git a/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs b/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs index 0d9cfcfd4cd..419430250c4 100644 --- a/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs @@ -1,6 +1,6 @@ use litellm_types::responses::streaming_websocket::{ResponsesWsEvent, ResponsesWsTransformResult}; -use crate::base_llm::chat::transformation::Error; +use crate::Error; pub const OPENAI_RESPONSES_DEFAULT_API_BASE: &str = "https://api.openai.com/v1"; pub const OPENAI_RESPONSES_PATH: &str = "/responses"; diff --git a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs index cfabcb12341..3525fd6322b 100644 --- a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs +++ b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs @@ -1,17 +1,20 @@ use litellm_auth_aws::{ - bedrock_model_id_and_region, + AwsCredentialSource, bedrock_model_id_and_region, constants::{BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE}, resolve_bedrock_region, }; use litellm_core_utils::core_helpers::json_type_name; use serde_json::{Map, Value, json}; -use crate::base_llm::{ - audio_transcription::transformation::{ - AudioTranscriptionRequestData, AudioTranscriptionResponseData, - BaseAudioTranscriptionConfig, RequestAuth, +use crate::{ + Error, + base_llm::{ + audio_transcription::transformation::{ + AudioTranscriptionRequestData, AudioTranscriptionResponseData, + BaseAudioTranscriptionConfig, Headers, ValidatedEnvironment, + }, + auth::AuthScheme, }, - chat::transformation::Error, }; const SUPPORTED_PARAMS: &[&str] = &["language", "prompt", "temperature", "response_format"]; @@ -131,16 +134,28 @@ impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig { )) } - fn auth_strategy( + fn validate_environment( &self, + headers: Headers, model: &str, optional_params: &Map, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { let (_, model_region) = bedrock_model_id_and_region(model); - Ok(RequestAuth::AwsSigV4 { - region: resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup), - service: BEDROCK_SERVICE, + Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::AwsSigV4 { + region: resolve_bedrock_region( + model_region.as_deref(), + optional_params, + env_lookup, + ), + service: BEDROCK_SERVICE, + credentials: Box::new(AwsCredentialSource::from_params( + optional_params, + env_lookup, + )), + }, }) } } diff --git a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs index b5db88d7dc4..8254165a739 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs @@ -1,5 +1,6 @@ +use litellm_auth::{CredentialPlacement, SecretValue}; use litellm_auth_aws::{ - bedrock_model_id_and_region, + AwsCredentialSource, bedrock_model_id_and_region, constants::{AWS_BEARER_TOKEN_BEDROCK, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE}, resolve_bedrock_region, }; @@ -16,9 +17,18 @@ use litellm_types::{ }; use serde_json::{Map, Value, json}; -use crate::base_llm::chat::transformation::{ - BaseConfig, Error, ProviderChatRequestData, ProviderChatResponseData, RequestAuth, Unsupported, - unsupported_message, unsupported_param, +use crate::{ + Error, + base_llm::{ + auth::AuthScheme, + chat::{ + streaming::StreamShape, + transformation::{ + BaseConfig, Headers, ProviderChatRequestData, ProviderChatResponseData, + Unsupported, ValidatedEnvironment, unsupported_message, unsupported_param, + }, + }, + }, }; /// Converse parameter names, post `map_openai_params`, that the Rust path can @@ -99,6 +109,7 @@ impl BaseConfig for AmazonConverseConfig { ) -> Result { Ok(ProviderChatRequestData { body: converse_body(&build_conversation(&messages), &optional_params), + stream_shape: StreamShape::default(), }) } @@ -180,31 +191,48 @@ impl BaseConfig for AmazonConverseConfig { }) } - fn auth( + /// Python reads `api_key` as the Bedrock bearer token and consults the env only when + /// the caller passed none, so a caller-supplied empty key falls through to SigV4 + /// without reaching for the environment. An all-whitespace token stays a bearer token + /// here because Python sends it too: treating it as absent would sign as the host + /// principal instead, which is the identity swap this branch exists to prevent. + fn validate_environment( &self, + headers: Headers, api_key: Option<&str>, model: &str, optional_params: &Map, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - // Python reads `api_key` as the Bedrock bearer token and consults the - // env only when the caller passed none, so a caller-supplied empty key - // falls through to SigV4 without reaching for the environment. An - // all-whitespace token stays a bearer token here because Python sends - // it too: treating it as absent would sign as the host principal - // instead, which is the identity swap this branch exists to prevent. + ) -> Result { let bearer = match api_key { Some(key) => Some(key.to_string()), None => env_lookup(AWS_BEARER_TOKEN_BEDROCK), } .filter(|token| !token.is_empty()); if let Some(token) = bearer { - return Ok(RequestAuth::Bearer { token }); + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret: SecretValue::new(token), + }, + }); } let (_, model_region) = bedrock_model_id_and_region(model); - Ok(RequestAuth::AwsSigV4 { - region: resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup), - service: BEDROCK_SERVICE, + Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::AwsSigV4 { + region: resolve_bedrock_region( + model_region.as_deref(), + optional_params, + env_lookup, + ), + service: BEDROCK_SERVICE, + credentials: Box::new(AwsCredentialSource::from_params( + optional_params, + env_lookup, + )), + }, }) } diff --git a/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs b/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs new file mode 100644 index 00000000000..b424dbd358d --- /dev/null +++ b/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs @@ -0,0 +1,138 @@ +use base64::Engine; +use bytes::Buf; +use futures_util::{Stream, StreamExt}; +use litellm_framing::{ + aws_event_stream::{AwsEventStreamCodec, Message}, + frames, +}; +use serde::Deserialize; +use serde_json::Value; + +use crate::{ + Error, + anthropic::{ + chat::handler::ModelResponseIterator, + messages::streaming_iterator::AnthropicMessagesStreamEvent, + }, + base_llm::{ + anthropic_messages::streaming::{ByteStream, EventStream}, + chat::streaming::{ChatStream, StreamShape}, + }, +}; + +#[derive(Deserialize)] +struct InvokeChunkPayload { + bytes: String, +} + +pub fn decode_invoke_chunk(message: Message) -> Result { + let payload: InvokeChunkPayload = + serde_json::from_slice(message.payload()).map_err(|error| { + Error::InvalidResponse(format!("Bedrock event payload is invalid: {error}")) + })?; + let chunk = base64::engine::general_purpose::STANDARD + .decode(payload.bytes) + .map_err(|error| { + Error::InvalidResponse(format!("Bedrock event payload has invalid base64: {error}")) + })?; + serde_json::from_slice(&chunk).map_err(|error| { + Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}")) + }) +} + +pub fn invoke_chunk_stream(input: S) -> impl Stream> + Send +where + S: Stream> + Send, + B: Buf + Send, + E: std::error::Error + Send + Sync + 'static, +{ + frames(input, AwsEventStreamCodec).map(|message| { + decode_invoke_chunk( + message.map_err(|error| { + Error::InvalidResponse(format!("stream framing failed: {error}")) + })?, + ) + }) +} + +pub fn decode_invoke_anthropic_chunk(chunk: Value) -> Result { + serde_json::from_value(chunk).map_err(|error| { + Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}")) + }) +} + +pub fn invoke_anthropic_event_stream(bytes: ByteStream) -> EventStream { + Box::pin(invoke_chunk_stream(bytes).map(|chunk| decode_invoke_anthropic_chunk(chunk?))) +} + +pub fn invoke_chat_stream(invoke_provider: &str, shape: StreamShape) -> Result { + match invoke_provider { + "anthropic" => Ok(ChatStream::new( + invoke_anthropic_event_stream, + ModelResponseIterator::new(shape), + )), + "deepseek_r1" | "moonshot" => Err(Error::Unsupported( + "Bedrock invoke streaming for this model family", + )), + _ => Err(Error::Unsupported("Bedrock invoke streaming")), + } +} + +#[cfg(test)] +mod tests { + use aws_smithy_eventstream::frame::write_message_to; + use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; + use base64::engine::general_purpose::STANDARD; + use bytes::Bytes; + use futures_util::TryStreamExt; + + use super::*; + use crate::{ + anthropic::messages::streaming_iterator::AnthropicContentBlockDelta, + base_llm::anthropic_messages::streaming::anthropic_sse_event_stream, + }; + + fn in_pieces(wire: &[u8]) -> ByteStream { + let pieces: Vec = wire.chunks(3).map(Bytes::copy_from_slice).collect(); + futures_util::stream::iter(pieces.into_iter().map(Ok)).boxed() + } + + const TEXT_DELTA: &str = + r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hello"}}"#; + + fn aws_wire(chunk: &str) -> Vec { + let payload = serde_json::json!({"bytes": STANDARD.encode(chunk)}); + let message = Message::new(Bytes::from(serde_json::to_vec(&payload).unwrap())).add_header( + Header::new(":event-type", HeaderValue::String("chunk".into())), + ); + let mut wire = Vec::new(); + write_message_to(&message, &mut wire).unwrap(); + wire + } + + #[tokio::test] + async fn aws_and_sse_framing_decode_to_the_same_anthropic_events() { + let aws = aws_wire(TEXT_DELTA); + let sse = format!("event: content_block_delta\ndata: {TEXT_DELTA}\n\n"); + + let from_aws = invoke_anthropic_event_stream(in_pieces(&aws)) + .try_collect::>() + .await + .unwrap(); + let from_sse = anthropic_sse_event_stream(in_pieces(sse.as_bytes())) + .try_collect::>() + .await + .unwrap(); + + assert_eq!( + from_aws, + vec![AnthropicMessagesStreamEvent::ContentBlockDelta { + index: 0, + delta: AnthropicContentBlockDelta::TextDelta { + text: "hello".into(), + }, + }] + ); + assert_eq!(from_aws, from_sse); + } +} diff --git a/litellm-rust/crates/llms/src/bedrock/chat/mod.rs b/litellm-rust/crates/llms/src/bedrock/chat/mod.rs index a41ad86ef49..a46514aa697 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/mod.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/mod.rs @@ -1 +1,2 @@ pub mod converse_transformation; +pub mod invoke_handler; diff --git a/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs new file mode 100644 index 00000000000..f2e365e9ed0 --- /dev/null +++ b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs @@ -0,0 +1,554 @@ +use std::convert::Infallible; + +use futures_util::StreamExt; +use litellm_auth::{CredentialPlacement, SecretValue}; +use litellm_auth_aws::{ + AwsCredentialSource, bedrock_model_id_and_region, + constants::{ + AWS_BEARER_TOKEN_BEDROCK, AWS_BEDROCK_RUNTIME_ENDPOINT, AWS_DEFAULT_REGION, AWS_REGION, + AWS_REGION_NAME, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE, + }, + resolve_bedrock_region, +}; +use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; +use serde_json::{Map, Value}; + +use crate::{ + Error, + anthropic::messages::streaming_iterator::{AnthropicMessagesStreamEvent, AnthropicStreamUsage}, + base_llm::{ + anthropic_messages::{ + streaming::{ByteStream, EventStream, StreamDecoder}, + transformation::{ + BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, + ValidatedEnvironment, + }, + }, + auth::AuthScheme, + base_model_iterator::{StreamError, StreamTransformer, transform_stream}, + }, + bedrock::chat::invoke_handler::{decode_invoke_anthropic_chunk, invoke_chunk_stream}, +}; + +const INVOCATION_METRICS_KEY: &str = "amazon-bedrock-invocationMetrics"; + +const METRICS_USAGE_KEYS: [(&str, &str); 4] = [ + ("input_tokens", "inputTokenCount"), + ("output_tokens", "outputTokenCount"), + ("cache_read_input_tokens", "cacheReadInputTokenCount"), + ("cache_creation_input_tokens", "cacheWriteInputTokenCount"), +]; + +const INVOKE_PATH: &str = "invoke"; +const INVOKE_STREAM_PATH: &str = "invoke-with-response-stream"; +const INVOKE_MODEL_PREFIX: &str = "invoke/"; + +const SECRET_NAMES: &[&str] = &[ + AWS_BEARER_TOKEN_BEDROCK, + AWS_BEDROCK_RUNTIME_ENDPOINT, + AWS_REGION_NAME, + AWS_REGION, + AWS_DEFAULT_REGION, +]; + +pub struct AmazonAnthropicClaudeMessagesConfig; + +pub const BEDROCK_ANTHROPIC_MESSAGES_CONFIG: AmazonAnthropicClaudeMessagesConfig = + AmazonAnthropicClaudeMessagesConfig; + +fn bearer_token( + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> Option { + match api_key { + Some(key) => Some(key.to_string()), + None => env_lookup(AWS_BEARER_TOKEN_BEDROCK), + } + .filter(|token| !token.is_empty()) +} + +fn invoke_url( + api_base: Option<&str>, + model: &str, + env_lookup: &dyn Fn(&str) -> Option, + path: &str, +) -> String { + let (model_id, model_region) = + bedrock_model_id_and_region(model.strip_prefix(INVOKE_MODEL_PREFIX).unwrap_or(model)); + let region = resolve_bedrock_region(model_region.as_deref(), &Map::new(), env_lookup); + let endpoint = api_base + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .or_else(|| env_lookup(AWS_BEDROCK_RUNTIME_ENDPOINT)) + .unwrap_or_else(|| BEDROCK_RUNTIME_ENDPOINT_TEMPLATE.replace("{region}", ®ion)); + format!("{}/model/{model_id}/{path}", endpoint.trim_end_matches('/')) +} + +impl BaseAnthropicMessagesConfig for AmazonAnthropicClaudeMessagesConfig { + fn get_complete_url( + &self, + api_base: Option<&str>, + model: &str, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + Ok(invoke_url(api_base, model, env_lookup, INVOKE_PATH)) + } + + fn complete_stream_url( + &self, + api_base: Option<&str>, + model: &str, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + Ok(invoke_url(api_base, model, env_lookup, INVOKE_STREAM_PATH)) + } + + fn transform_anthropic_messages_request( + &self, + _request: AnthropicMessagesRequest, + _context: &MessagesTransformContext, + ) -> Result { + Err(Error::Unsupported( + "Bedrock invoke messages request shaping", + )) + } + + fn secret_names(&self) -> &'static [&'static str] { + SECRET_NAMES + } + + /// Python reads `api_key` as the Bedrock bearer token and consults the env only when the + /// caller passed none. Without one the request is signed with SigV4. + fn validate_environment( + &self, + headers: Headers, + api_key: Option<&str>, + model: &str, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + if let Some(token) = bearer_token(api_key, env_lookup) { + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret: SecretValue::new(token), + }, + }); + } + let (_, model_region) = + bedrock_model_id_and_region(model.strip_prefix(INVOKE_MODEL_PREFIX).unwrap_or(model)); + let params = Map::new(); + Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::AwsSigV4 { + region: resolve_bedrock_region(model_region.as_deref(), ¶ms, env_lookup), + service: BEDROCK_SERVICE, + credentials: Box::new(AwsCredentialSource::from_params(¶ms, env_lookup)), + }, + }) + } + + fn default_headers(&self) -> &'static [(&'static str, &'static str)] { + &[("content-type", "application/json")] + } + + fn stream_decoder(&self) -> Option { + Some(bedrock_anthropic_messages_event_stream) + } +} + +fn with_invocation_usage(chunk: Value) -> Value { + match chunk { + Value::Object(fields) => Value::Object(with_metrics_usage(fields)), + other => other, + } +} + +fn with_metrics_usage(mut fields: Map) -> Map { + let Some(Value::Object(metrics)) = fields.remove(INVOCATION_METRICS_KEY) else { + return fields; + }; + if metrics.is_empty() { + return fields; + } + let preserved = match fields.remove("usage") { + Some(Value::Object(usage)) => usage, + _ => Map::new(), + }; + let usage: Map = METRICS_USAGE_KEYS + .iter() + .filter_map(|(anthropic, metric)| { + Some((anthropic.to_string(), metrics.get(*metric)?.clone())) + }) + .chain(preserved) + .collect(); + fields.insert("usage".to_string(), Value::Object(usage)); + fields +} + +pub fn bedrock_anthropic_messages_event_stream(bytes: ByteStream) -> EventStream { + let events = invoke_chunk_stream(bytes) + .map(|chunk| decode_invoke_anthropic_chunk(with_invocation_usage(chunk?))); + Box::pin( + transform_stream(events, MessageStopUsagePromoter::default()) + .map(|item| item.map_err(StreamError::into_decode)), + ) +} + +#[derive(Default)] +pub struct MessageStopUsagePromoter { + pending_delta: Option, + start_usage: Option, +} + +fn promoted_usage( + delta: Option, + stop: Option<&AnthropicStreamUsage>, + start: Option<&AnthropicStreamUsage>, +) -> Option { + let delta = delta.unwrap_or_default(); + let merged = AnthropicStreamUsage { + input_tokens: stop + .and_then(|stop| stop.input_tokens) + .or(delta.input_tokens), + cache_creation_input_tokens: stop + .and_then(|stop| stop.cache_creation_input_tokens) + .or(delta.cache_creation_input_tokens) + .or_else(|| start.and_then(|start| start.cache_creation_input_tokens)), + cache_read_input_tokens: stop + .and_then(|stop| stop.cache_read_input_tokens) + .or(delta.cache_read_input_tokens) + .or_else(|| start.and_then(|start| start.cache_read_input_tokens)), + extra: delta + .extra + .into_iter() + .chain( + start + .and_then(|start| start.extra.get_key_value("cache_creation")) + .map(|(key, value)| (key.clone(), value.clone())), + ) + .fold(Map::new(), |mut extra, (key, value)| { + extra.entry(key).or_insert(value); + extra + }), + ..delta + }; + (merged != AnthropicStreamUsage::default()).then_some(merged) +} + +fn promoted( + event: AnthropicMessagesStreamEvent, + stop: Option<&AnthropicStreamUsage>, + start: Option<&AnthropicStreamUsage>, +) -> AnthropicMessagesStreamEvent { + match event { + AnthropicMessagesStreamEvent::MessageDelta { + delta, + usage, + context_management, + } => AnthropicMessagesStreamEvent::MessageDelta { + delta, + usage: promoted_usage(usage, stop, start), + context_management, + }, + other => other, + } +} + +impl StreamTransformer for MessageStopUsagePromoter { + type Input = AnthropicMessagesStreamEvent; + type Output = AnthropicMessagesStreamEvent; + type Error = Infallible; + + fn transform( + &mut self, + input: AnthropicMessagesStreamEvent, + ) -> Result, Infallible> { + let pending = self.pending_delta.take(); + match input { + AnthropicMessagesStreamEvent::MessageDelta { .. } => { + self.pending_delta = Some(input); + Ok(pending.into_iter().collect()) + } + AnthropicMessagesStreamEvent::MessageStop { usage } => Ok(pending + .map(|delta| promoted(delta, usage.as_ref(), self.start_usage.as_ref())) + .into_iter() + .chain([AnthropicMessagesStreamEvent::MessageStop { usage }]) + .collect()), + AnthropicMessagesStreamEvent::MessageStart { message } => { + self.start_usage = Some(message.usage.clone()); + Ok(pending + .into_iter() + .chain([AnthropicMessagesStreamEvent::MessageStart { message }]) + .collect()) + } + other => Ok(pending.into_iter().chain([other]).collect()), + } + } + + fn finish(&mut self) -> Result, Infallible> { + Ok(self + .pending_delta + .take() + .map(|delta| promoted(delta, None, self.start_usage.as_ref())) + .into_iter() + .collect()) + } +} + +#[cfg(test)] +mod tests { + use aws_smithy_eventstream::frame::write_message_to; + use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; + use base64::{Engine, engine::general_purpose::STANDARD}; + use bytes::Bytes; + use futures_util::TryStreamExt; + use rstest::rstest; + use serde_json::json; + + use litellm_auth_aws::constants::DEFAULT_BEDROCK_REGION; + + use super::*; + use crate::base_llm::anthropic_messages::streaming::encode_anthropic_sse; + + fn event(value: Value) -> AnthropicMessagesStreamEvent { + serde_json::from_value(value).unwrap() + } + + fn message_start(usage: Value) -> AnthropicMessagesStreamEvent { + event(json!({ + "type": "message_start", + "message": { + "id": "msg_1", "type": "message", "role": "assistant", "model": "m", + "content": [], "stop_reason": null, "stop_sequence": null, "usage": usage + } + })) + } + + fn message_delta(usage: Value) -> AnthropicMessagesStreamEvent { + event(json!({ + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + "usage": usage + })) + } + + fn message_stop(usage: Option) -> AnthropicMessagesStreamEvent { + match usage { + Some(usage) => event(json!({"type": "message_stop", "usage": usage})), + None => event(json!({"type": "message_stop"})), + } + } + + fn promote(events: Vec) -> Vec { + let mut promoter = MessageStopUsagePromoter::default(); + let mut output: Vec<_> = events + .into_iter() + .flat_map(|event| promoter.transform(event).unwrap()) + .collect(); + output.extend(promoter.finish().unwrap()); + output + } + + #[rstest] + #[case::cache_fields_on_message_stop( + json!({"input_tokens": 10, "output_tokens": 0}), + json!({"output_tokens": 5}), + Some(json!({"input_tokens": 3, "cache_read_input_tokens": 100, "cache_creation_input_tokens": 20})), + json!({"input_tokens": 3, "output_tokens": 5, "cache_read_input_tokens": 100, "cache_creation_input_tokens": 20}), + )] + #[case::cache_only_on_message_start( + json!({"input_tokens": 10, "output_tokens": 0, "cache_read_input_tokens": 80, "cache_creation_input_tokens": 4, "cache_creation": {"ephemeral_5m_input_tokens": 4}}), + json!({"output_tokens": 5}), + Some(json!({"input_tokens": 10})), + json!({"input_tokens": 10, "output_tokens": 5, "cache_read_input_tokens": 80, "cache_creation_input_tokens": 4, "cache_creation": {"ephemeral_5m_input_tokens": 4}}), + )] + #[case::message_stop_wins_over_message_start( + json!({"input_tokens": 10, "cache_read_input_tokens": 80}), + json!({"output_tokens": 5}), + Some(json!({"cache_read_input_tokens": 100})), + json!({"output_tokens": 5, "cache_read_input_tokens": 100}), + )] + #[case::delta_cache_fields_are_kept( + json!({"input_tokens": 10, "cache_read_input_tokens": 80}), + json!({"output_tokens": 5, "cache_read_input_tokens": 7}), + None, + json!({"output_tokens": 5, "cache_read_input_tokens": 7}), + )] + fn message_delta_usage_is_completed_from_stop_then_start( + #[case] start: Value, + #[case] delta: Value, + #[case] stop: Option, + #[case] expected: Value, + ) { + let output = promote(vec![ + message_start(start), + message_delta(delta), + message_stop(stop.clone()), + ]); + + assert_eq!(output.len(), 3); + assert_eq!(output[1], message_delta(expected)); + assert_eq!(output[2], message_stop(stop)); + } + + #[test] + fn a_delta_is_flushed_with_start_usage_when_the_stream_ends_without_a_stop() { + let output = promote(vec![ + message_start(json!({"input_tokens": 10, "cache_read_input_tokens": 80})), + message_delta(json!({"output_tokens": 5})), + ]); + + assert_eq!( + output[1], + message_delta(json!({"output_tokens": 5, "cache_read_input_tokens": 80})) + ); + } + + #[test] + fn events_around_the_delta_keep_their_order() { + let ping = event(json!({"type": "ping"})); + let output = promote(vec![ + message_delta(json!({"output_tokens": 5})), + ping.clone(), + message_stop(None), + ]); + + assert_eq!( + output, + vec![ + message_delta(json!({"output_tokens": 5})), + ping, + message_stop(None) + ] + ); + } + + #[rstest] + #[case::metrics_fill_missing_usage( + json!({"type": "message_stop", "amazon-bedrock-invocationMetrics": {"inputTokenCount": 3, "outputTokenCount": 9}}), + json!({"type": "message_stop", "usage": {"input_tokens": 3, "output_tokens": 9}}), + )] + #[case::the_chunks_own_usage_wins( + json!({"type": "message_stop", "usage": {"input_tokens": 1}, "amazon-bedrock-invocationMetrics": {"inputTokenCount": 3, "cacheReadInputTokenCount": 40}}), + json!({"type": "message_stop", "usage": {"cache_read_input_tokens": 40, "input_tokens": 1}}), + )] + #[case::no_metrics_leaves_the_chunk( + json!({"type": "message_stop"}), + json!({"type": "message_stop"}), + )] + #[case::empty_metrics_are_dropped( + json!({"type": "message_stop", "amazon-bedrock-invocationMetrics": {}}), + json!({"type": "message_stop"}), + )] + fn invocation_metrics_become_anthropic_usage(#[case] chunk: Value, #[case] expected: Value) { + assert_eq!(with_invocation_usage(chunk), expected); + } + + fn aws_frame(chunk: &Value) -> Vec { + let payload = json!({"bytes": STANDARD.encode(chunk.to_string())}); + let message = Message::new(Bytes::from(serde_json::to_vec(&payload).unwrap())).add_header( + Header::new(":event-type", HeaderValue::String("chunk".into())), + ); + let mut wire = Vec::new(); + write_message_to(&message, &mut wire).unwrap(); + wire + } + + #[tokio::test] + async fn bedrock_stream_yields_the_sse_an_anthropic_client_reads() { + let chunks = [ + json!({"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 5}}), + json!({"type": "message_stop", "amazon-bedrock-invocationMetrics": {"inputTokenCount": 3, "cacheReadInputTokenCount": 40}}), + ]; + let wire: Vec = chunks.iter().flat_map(aws_frame).collect(); + let bytes: ByteStream = futures_util::stream::iter( + wire.chunks(7) + .map(|chunk| Ok(Bytes::copy_from_slice(chunk))) + .collect::>(), + ) + .boxed(); + + let sse = bedrock_anthropic_messages_event_stream(bytes) + .map_ok(|event| encode_anthropic_sse(&event).unwrap()) + .try_collect::>() + .await + .unwrap() + .concat(); + + let expected: Vec = [ + message_delta( + json!({"output_tokens": 5, "cache_read_input_tokens": 40, "input_tokens": 3}), + ), + message_stop(Some( + json!({"input_tokens": 3, "cache_read_input_tokens": 40}), + )), + ] + .iter() + .flat_map(|event| encode_anthropic_sse(event).unwrap()) + .collect(); + assert_eq!(sse, expected); + } + + #[test] + fn config_uses_the_streaming_url_only_for_streams() { + let env = |_: &str| -> Option { None }; + let config = AmazonAnthropicClaudeMessagesConfig; + + assert_eq!( + config + .get_complete_url(None, "anthropic.claude-3", &env) + .unwrap(), + config + .complete_stream_url(None, "anthropic.claude-3", &env) + .unwrap() + .replace(INVOKE_STREAM_PATH, INVOKE_PATH) + ); + } + + #[rstest] + #[case::an_explicit_key_is_a_bearer_token(Some("token"), None, Some("token"))] + #[case::the_env_token_is_a_bearer_token(None, Some("env-token"), Some("env-token"))] + #[case::no_token_signs_with_sigv4(None, None, None)] + fn requests_sign_only_without_a_bearer_token( + #[case] api_key: Option<&str>, + #[case] env_token: Option<&str>, + #[case] expected_bearer: Option<&str>, + ) { + let env = |name: &str| { + (name == AWS_BEARER_TOKEN_BEDROCK) + .then(|| env_token.map(str::to_string)) + .flatten() + }; + let validated = AmazonAnthropicClaudeMessagesConfig + .validate_environment( + vec![("authorization".into(), "Bearer forwarded".into())], + api_key, + "anthropic.claude-3", + &env, + ) + .unwrap(); + match (validated.auth, expected_bearer) { + ( + AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret, + }, + Some(expected), + ) => assert_eq!(secret.expose(), expected), + ( + AuthScheme::AwsSigV4 { + region, service, .. + }, + None, + ) => { + assert_eq!( + (region.as_str(), service), + (DEFAULT_BEDROCK_REGION, BEDROCK_SERVICE) + ); + } + (other, _) => panic!("unexpected auth {other:?}"), + } + } +} diff --git a/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/mod.rs b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/mod.rs new file mode 100644 index 00000000000..4d67a0c0696 --- /dev/null +++ b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/mod.rs @@ -0,0 +1 @@ +pub mod anthropic_claude3_transformation; diff --git a/litellm-rust/crates/llms/src/bedrock/messages/mod.rs b/litellm-rust/crates/llms/src/bedrock/messages/mod.rs new file mode 100644 index 00000000000..476a99539ff --- /dev/null +++ b/litellm-rust/crates/llms/src/bedrock/messages/mod.rs @@ -0,0 +1 @@ +pub mod invoke_transformations; diff --git a/litellm-rust/crates/llms/src/bedrock/mod.rs b/litellm-rust/crates/llms/src/bedrock/mod.rs index 695aeb8af5e..feed6e70e4d 100644 --- a/litellm-rust/crates/llms/src/bedrock/mod.rs +++ b/litellm-rust/crates/llms/src/bedrock/mod.rs @@ -1,2 +1,3 @@ pub mod audio_transcription; pub mod chat; +pub mod messages; diff --git a/litellm-rust/crates/llms/src/error.rs b/litellm-rust/crates/llms/src/error.rs new file mode 100644 index 00000000000..e885d6f43a1 --- /dev/null +++ b/litellm-rust/crates/llms/src/error.rs @@ -0,0 +1,18 @@ +#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] +pub enum Error { + #[error("expected {expected}, got {actual}")] + InvalidType { + expected: &'static str, + actual: &'static str, + }, + #[error("missing required field: {0}")] + MissingField(&'static str), + #[error("invalid request: {0}")] + InvalidRequest(String), + #[error("invalid response: {0}")] + InvalidResponse(String), + #[error("unsupported: {0}")] + Unsupported(&'static str), + #[error(transparent)] + Auth(#[from] litellm_auth::Error), +} diff --git a/litellm-rust/crates/llms/src/lib.rs b/litellm-rust/crates/llms/src/lib.rs index 701eaff4374..e71a9466c0c 100644 --- a/litellm-rust/crates/llms/src/lib.rs +++ b/litellm-rust/crates/llms/src/lib.rs @@ -4,7 +4,10 @@ pub mod azure_ai; pub mod base_llm; pub mod bedrock; pub mod cohere; +mod error; pub mod mistral; pub mod openai; pub mod reducto; pub mod vertex_ai; + +pub use error::Error; diff --git a/litellm-rust/crates/llms/src/openai/responses/transformation.rs b/litellm-rust/crates/llms/src/openai/responses/transformation.rs index f01ec4ad146..1001265413e 100644 --- a/litellm-rust/crates/llms/src/openai/responses/transformation.rs +++ b/litellm-rust/crates/llms/src/openai/responses/transformation.rs @@ -1,8 +1,8 @@ use litellm_types::responses::streaming_websocket::{ResponsesWsEvent, ResponsesWsTransformResult}; -use crate::base_llm::{ - chat::transformation::Error, - responses::transformation::{ResponsesWebSocketProviderConfig, enforce_model}, +use crate::{ + Error, + base_llm::responses::transformation::{ResponsesWebSocketProviderConfig, enforce_model}, }; pub struct OpenAiResponsesApiConfig; diff --git a/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs b/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs index c3377536545..147056dab8d 100644 --- a/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs @@ -564,7 +564,10 @@ mod tests { let params = ReductoParseV3Config .map_ocr_params(&overrides, "parse-v3") .unwrap(); - let client = OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new()); + let client = OcrClient::for_test( + litellm_http::Client::plain_for_test(), + litellm_http::Client::no_redirect_for_test(), + ); let connection = OcrConnection::default(); let document = serde_json::from_value( json!({"type":"document_url","document_url":"reducto://ready.pdf"}), diff --git a/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs index c9342c87e9a..58e2f6cb0ad 100644 --- a/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs @@ -130,7 +130,8 @@ impl VertexAiOcrConfig { ) -> Result { validate_destination(connection)?; client - .vertex_auth() + .auth() + .gcp .validate_environment( connection.extra_headers.clone(), connection diff --git a/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs b/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs index ed22a1d141d..f6f0b8eed42 100644 --- a/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs +++ b/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs @@ -1,7 +1,9 @@ use litellm_llms::{ + Error, anthropic::chat::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG, - base_llm::chat::transformation::{ - BaseConfig, Error, ProviderChatResponseData, RequestAuth, Unsupported, + base_llm::{ + auth::AuthScheme, + chat::transformation::{BaseConfig, ProviderChatResponseData, Unsupported}, }, }; use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse}; @@ -430,15 +432,22 @@ fn resolves_the_messages_url_and_x_api_key_auth() { .expect("url builds"), "https://api.anthropic.com/v1/messages" ); - assert_eq!( - config - .auth(Some("sk-x"), "claude-sonnet-4-5", &Map::new(), &|_| None) - .expect("auth resolves"), - RequestAuth::Header { - name: "x-api-key", - value: "sk-x".to_string() - } - ); + let validated = config + .validate_environment( + Vec::new(), + Some("sk-x"), + "claude-sonnet-4-5", + &Map::new(), + &|_| None, + ) + .expect("auth resolves"); + assert!(matches!( + validated.auth, + AuthScheme::Credential { + placement: litellm_auth::CredentialPlacement::Header("x-api-key"), + ref secret + } if secret.expose() == "sk-x" + )); assert_eq!( config.default_headers(), &[ diff --git a/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs b/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs index 4127bcfa19d..0bd637f602c 100644 --- a/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs +++ b/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs @@ -1,6 +1,9 @@ +use litellm_auth::CredentialPlacement; use litellm_llms::{ - base_llm::chat::transformation::{ - BaseConfig, Error, ProviderChatResponseData, RequestAuth, Unsupported, + Error, + base_llm::{ + auth::AuthScheme, + chat::transformation::{BaseConfig, ProviderChatResponseData, Unsupported}, }, bedrock::chat::converse_transformation::BEDROCK_CHAT_COMPLETIONS_CONFIG, }; @@ -273,22 +276,36 @@ fn prefers_an_explicit_runtime_endpoint_over_the_api_base() { ); } +/// The bearer token a config named, or `None` for a SigV4 scheme in the given region. +fn bearer_or_region(auth: AuthScheme) -> Result { + match auth { + AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret, + } => Ok(secret.expose().to_string()), + AuthScheme::AwsSigV4 { + region, + service: "bedrock", + .. + } => Err(region), + other => panic!("unexpected auth {other:?}"), + } +} + #[test] fn signs_with_sigv4_in_the_resolved_region() { - let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG; + let validated = BEDROCK_CHAT_COMPLETIONS_CONFIG + .validate_environment( + Vec::new(), + None, + "eu-central-1/anthropic.claude-v2", + &Map::new(), + &|_| None, + ) + .expect("auth resolves"); assert_eq!( - config - .auth( - None, - "eu-central-1/anthropic.claude-v2", - &Map::new(), - &|_| None - ) - .expect("auth resolves"), - RequestAuth::AwsSigV4 { - region: "eu-central-1".to_string(), - service: "bedrock", - } + bearer_or_region(validated.auth), + Err("eu-central-1".to_string()) ); } @@ -302,22 +319,21 @@ fn a_bearer_token_outranks_sigv4_the_way_python_resolves_it() { |key: &str| (key == "AWS_BEARER_TOKEN_BEDROCK").then(|| "from-env".to_string()); let no_env = |_: &str| None; let resolve = |api_key, env: &dyn Fn(&str) -> Option| { - BEDROCK_CHAT_COMPLETIONS_CONFIG - .auth( - api_key, - "eu-central-1/anthropic.claude-v2", - &Map::new(), - env, - ) - .expect("auth resolves") - }; - let bearer = |token: &str| RequestAuth::Bearer { - token: token.to_string(), - }; - let sigv4 = RequestAuth::AwsSigV4 { - region: "eu-central-1".to_string(), - service: "bedrock", + bearer_or_region( + BEDROCK_CHAT_COMPLETIONS_CONFIG + .validate_environment( + Vec::new(), + api_key, + "eu-central-1/anthropic.claude-v2", + &Map::new(), + env, + ) + .expect("auth resolves") + .auth, + ) }; + let bearer = |token: &str| Ok(token.to_string()); + let sigv4 = Err("eu-central-1".to_string()); // A caller-supplied key is the bearer token, and outranks the env. assert_eq!( diff --git a/litellm-rust/crates/llms/tests/ocr_handler.rs b/litellm-rust/crates/llms/tests/ocr_handler.rs index 6e46e6f76d4..5ed3244087d 100644 --- a/litellm-rust/crates/llms/tests/ocr_handler.rs +++ b/litellm-rust/crates/llms/tests/ocr_handler.rs @@ -19,7 +19,7 @@ async fn read_bounded(response: String, limit: usize) -> Result().await; }); - let response = reqwest::Client::new() + let response = litellm_http::Client::plain_for_test() .get(format!("http://{address}")) .send() .await diff --git a/litellm-rust/crates/model-catalog/Cargo.toml b/litellm-rust/crates/model-catalog/Cargo.toml index 0b26e398ac8..94a69c94fdf 100644 --- a/litellm-rust/crates/model-catalog/Cargo.toml +++ b/litellm-rust/crates/model-catalog/Cargo.toml @@ -6,11 +6,13 @@ license.workspace = true repository.workspace = true [features] -schema = ["dep:schemars"] +schema = ["dep:schemars", "litellm-types/schema"] [dependencies] +litellm-types.workspace = true + indexmap = { version = "2.14.0", features = ["serde"] } -schemars = { version = "1.0", optional = true } +schemars = { workspace = true, optional = true } serde.workspace = true serde_json.workspace = true thiserror.workspace = true diff --git a/litellm-rust/crates/model-catalog/src/capabilities.rs b/litellm-rust/crates/model-catalog/src/capabilities.rs index 66b5f1c5d2e..3df68fabc4d 100644 --- a/litellm-rust/crates/model-catalog/src/capabilities.rs +++ b/litellm-rust/crates/model-catalog/src/capabilities.rs @@ -24,20 +24,6 @@ pub enum Mode { VideoGeneration, } -/// Reasoning effort level accepted or applied by the model. -#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(rename_all = "snake_case")] -pub enum ReasoningEffort { - None, - Minimal, - Low, - Medium, - High, - Xhigh, - Max, -} - /// Gemini audio generation API the model is served through. #[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] #[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] diff --git a/litellm-rust/crates/model-catalog/src/model_info.rs b/litellm-rust/crates/model-catalog/src/model_info.rs index 361cb56e9b1..7b8ce15fbd6 100644 --- a/litellm-rust/crates/model-catalog/src/model_info.rs +++ b/litellm-rust/crates/model-catalog/src/model_info.rs @@ -1,7 +1,6 @@ -use crate::capabilities::{ - AudioFormat, InputModality, Mode, OutputModality, ReasoningEffort, VertexAiAudioApi, -}; +use crate::capabilities::{AudioFormat, InputModality, Mode, OutputModality, VertexAiAudioApi}; use crate::pricing::{OffPeakPricing, SearchContextCostPerQuery, TieredRate, WebSearchBillingUnit}; +use litellm_types::llms::openai::ReasoningEffort; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::collections::BTreeMap; @@ -42,6 +41,9 @@ pub struct ModelInfo { pub cache_creation_input_token_cost_above_200k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] + pub cache_creation_input_token_cost_above_200k_tokens_batches: Option, + /// Rate applied once the prompt exceeds the token threshold in the field name. + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_256k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] @@ -78,6 +80,9 @@ pub struct ModelInfo { /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_200k_tokens: Option, + /// Rate applied once the prompt exceeds the token threshold in the field name. + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_read_input_token_cost_above_200k_tokens_batches: Option, /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_200k_tokens_priority: Option, @@ -113,6 +118,10 @@ pub struct ModelInfo { pub code_interpreter_cost_per_session: Option, #[serde(skip_serializing_if = "Option::is_none")] pub comment: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub computer_use_input_cost_per_1k_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub computer_use_output_cost_per_1k_tokens: Option, /// Reasoning effort the provider applies when the request omits reasoning_effort. Gates whether a non-default temperature or the top_p/logprobs sampling params are accepted, which hold only when the effort resolves to 'none'. #[serde(skip_serializing_if = "Option::is_none")] pub default_reasoning_effort: Option, @@ -120,6 +129,10 @@ pub struct ModelInfo { #[serde(skip_serializing_if = "Option::is_none")] pub deprecation_date: Option, #[serde(skip_serializing_if = "Option::is_none")] + pub file_search_cost_per_1k_calls: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub file_search_cost_per_gb_per_day: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub gemini_audio_only_live: Option, #[serde(skip_serializing_if = "Option::is_none")] pub gemini_native_audio: Option, @@ -174,6 +187,9 @@ pub struct ModelInfo { /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_200k_tokens: Option, + /// Rate applied once the prompt exceeds the token threshold in the field name. + #[serde(skip_serializing_if = "Option::is_none")] + pub input_cost_per_token_above_200k_tokens_batches: Option, /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_200k_tokens_priority: Option, @@ -265,6 +281,26 @@ pub struct ModelInfo { pub output_cost_per_image_1536: Option, #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image_512: Option, + #[serde( + rename = "output_cost_per_image_0.5K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_0_5k: Option, + #[serde( + rename = "output_cost_per_image_1K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_1k: Option, + #[serde( + rename = "output_cost_per_image_2K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_2k: Option, + #[serde( + rename = "output_cost_per_image_4K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_4k: Option, #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image_token: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -297,6 +333,9 @@ pub struct ModelInfo { /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_200k_tokens: Option, + /// Rate applied once the prompt exceeds the token threshold in the field name. + #[serde(skip_serializing_if = "Option::is_none")] + pub output_cost_per_token_above_200k_tokens_batches: Option, /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_200k_tokens_priority: Option, @@ -357,6 +396,8 @@ pub struct ModelInfo { /// Provider default requests-per-minute limit. #[serde(skip_serializing_if = "Option::is_none")] pub rpm: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub rules: Option>, /// USD cost per web search query, keyed by search context size. #[serde(skip_serializing_if = "Option::is_none")] pub search_context_cost_per_query: Option, @@ -475,6 +516,8 @@ pub struct ModelInfo { #[serde(skip_serializing_if = "Option::is_none")] pub uses_embed_content: Option, #[serde(skip_serializing_if = "Option::is_none")] + pub vector_store_cost_per_gb_per_day: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub vertex_ai_audio_api: Option, /// Whether web search is billed per query or per prompt. #[serde(skip_serializing_if = "Option::is_none")] diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index a02adfaa064..7cfb3f207d4 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -42,7 +42,6 @@ litellm-auth-aws.workspace = true litellm-callbacks-legacy-python.workspace = true litellm-core.workspace = true litellm-core-utils.workspace = true -litellm-auth-gcp.workspace = true litellm-http.workspace = true litellm-llms.workspace = true litellm-secrets = { workspace = true, features = ["aws", "azure", "google", "hashicorp", "cyberark"] } @@ -55,12 +54,14 @@ pyo3-async-runtimes.workspace = true reqwest.workspace = true redis = { version = "1.7.0", features = ["tls-rustls"] } serde_json.workspace = true +strum.workspace = true veil.workspace = true thiserror.workspace = true tokio = { workspace = true, features = ["rt", "sync"] } url.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } litellm-secrets-aws.workspace = true serde.workspace = true serde_with.workspace = true diff --git a/litellm-rust/crates/python-bridge/preflight_contract.json b/litellm-rust/crates/python-bridge/preflight_contract.json new file mode 100644 index 00000000000..343dea268cd --- /dev/null +++ b/litellm-rust/crates/python-bridge/preflight_contract.json @@ -0,0 +1,10 @@ +{ + "credential_list": [], + "warn_unknown_credential": [ + "name", + "loaded" + ], + "check_limits": [ + "kwargs" + ] +} diff --git a/litellm-rust/crates/python-bridge/src/cache/activation.rs b/litellm-rust/crates/python-bridge/src/cache/activation.rs index f77032c579d..58735679554 100644 --- a/litellm-rust/crates/python-bridge/src/cache/activation.rs +++ b/litellm-rust/crates/python-bridge/src/cache/activation.rs @@ -9,10 +9,10 @@ use super::{ cache_error, config::{CacheBackendConfig, NativeCacheConfig, UnsupportedCacheConfig}, embedder::PythonEmbedder, - host_client, native::NativeResponseCache, }; use crate::errors::RustBridgeDeclined; +use crate::http::host_client; fn declined(reason: UnsupportedCacheConfig) -> PyErr { RustBridgeDeclined::new_err(reason.message()) diff --git a/litellm-rust/crates/python-bridge/src/cache/config.rs b/litellm-rust/crates/python-bridge/src/cache/config.rs index 6e25f07efa1..e58902b07ee 100644 --- a/litellm-rust/crates/python-bridge/src/cache/config.rs +++ b/litellm-rust/crates/python-bridge/src/cache/config.rs @@ -1511,7 +1511,7 @@ mod tests { path_service_account: Some("credentials.json".into()), endpoint: litellm_cache_gcs::DEFAULT_ENDPOINT.into(), }, - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), Some("token".into()), ); let matching_config = NativeCacheConfig { diff --git a/litellm-rust/crates/python-bridge/src/cache/handle.rs b/litellm-rust/crates/python-bridge/src/cache/handle.rs index e14916b25c6..fcc8aa6218a 100644 --- a/litellm-rust/crates/python-bridge/src/cache/handle.rs +++ b/litellm-rust/crates/python-bridge/src/cache/handle.rs @@ -1,3 +1,4 @@ +use crate::http::host_client; use crate::logger::run_sync_value; use litellm_auth_aws::AwsAuthConfig; use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig}; @@ -19,7 +20,6 @@ use super::{ config::{QdrantSemanticCacheConfig, project_redis_semantic}, embedder::PythonEmbedder, facade::FacadeGuard, - host_client, native::NativeResponseCache, request::duration, }; diff --git a/litellm-rust/crates/python-bridge/src/cache/mod.rs b/litellm-rust/crates/python-bridge/src/cache/mod.rs index ac1e00d5273..00b0c71684a 100644 --- a/litellm-rust/crates/python-bridge/src/cache/mod.rs +++ b/litellm-rust/crates/python-bridge/src/cache/mod.rs @@ -13,11 +13,9 @@ mod resolver; mod semantic; use litellm_cache::Error; -use litellm_http::ClientVariant; use pyo3::{ exceptions::{PyNotImplementedError, PyRuntimeError, PyValueError}, prelude::*, - types::PyDict, }; pub(crate) use self::{binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheResolver}; @@ -29,11 +27,3 @@ fn cache_error(error: Error) -> PyErr { _ => PyRuntimeError::new_err(error.to_string()), } } - -/// The host's pooled HTTP client, configured from the proxy's HTTP settings. -fn host_client(py: Python<'_>, variant: ClientVariant) -> PyResult { - let http_config = crate::http::call_config(py, &PyDict::new(py), true)?; - crate::http::pool() - .client(&http_config, variant) - .map_err(crate::http::client_error) -} diff --git a/litellm-rust/crates/python-bridge/src/cache/native.rs b/litellm-rust/crates/python-bridge/src/cache/native.rs index 0e279046812..460136baa1f 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native.rs @@ -92,7 +92,7 @@ impl NativeResponseCache { )) } - pub async fn s3(config: S3CacheConfig, http: reqwest::Client) -> Self { + pub async fn s3(config: S3CacheConfig, http: litellm_http::Client) -> Self { let runtime = tokio::runtime::Handle::current(); let backend = S3Cache::new(config, http, ResponseCacheCodec, runtime); let identity = BackendIdentity::S3 { @@ -112,7 +112,7 @@ impl NativeResponseCache { Ok(Self::exact(ResponseCache::new(Arc::new(backend)), identity)) } - pub fn gcs(config: GcsConfig, client: reqwest::Client, token: Option) -> Self { + pub fn gcs(config: GcsConfig, client: litellm_http::Client, token: Option) -> Self { let backend = match token { Some(token) => GcsCache::with_token_source( config, @@ -133,7 +133,7 @@ impl NativeResponseCache { pub async fn azure_blob( account_url: &str, container: &str, - http: reqwest::Client, + http: litellm_http::Client, ) -> Result { let backend = AzureBlobCache::connect( account_url, @@ -242,7 +242,7 @@ impl NativeResponseCache { pub async fn qdrant_semantic( config: QdrantSemanticCacheConfig, - client: reqwest::Client, + client: litellm_http::Client, runtime: tokio::runtime::Handle, ) -> Result { let qdrant = qdrant_client::Qdrant::from_url(&config.grpc_url) diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 6c5a65173e3..e91dd15beb0 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -1,6 +1,5 @@ -use litellm_core::{Error, audio_transcription, chat_completions, messages, responses}; +use litellm_core::{Phase, RouteError}; use litellm_http::transport::Error as TransportError; -use litellm_llms::base_llm::ocr::error::Error as OcrError; use pyo3::{ exceptions::{PyRuntimeError, PyValueError}, prelude::*, @@ -20,73 +19,16 @@ pyo3::create_exception!( "The provider call was already issued and failed. Args are (status, message); status is 0 when there was no HTTP response." ); -fn auth_is_value_error(error: &litellm_auth::Error) -> bool { - !matches!(error, litellm_auth::Error::MissingApiKey { .. }) +pub(crate) fn route_error_to_pyerr(error: RouteError) -> PyErr { + by_fault(error.is_request(), error.to_string()) } -pub(crate) fn messages_error_to_pyerr(error: messages::Error) -> PyErr { - core_error_to_pyerr(error.into()) -} - -pub(crate) fn audio_transcription_error_to_pyerr(error: audio_transcription::Error) -> PyErr { - core_error_to_pyerr(error.into()) -} - -pub(crate) fn responses_error_to_pyerr(error: responses::Error) -> PyErr { - core_error_to_pyerr(error.into()) -} - -pub(crate) fn core_error_to_pyerr(error: Error) -> PyErr { - let value_error = match &error { - Error::Ocr(error) => { - error.is_request() - || matches!( - error, - OcrError::Auth(_) - | OcrError::InvalidProvider(_) - | OcrError::InvalidRequest(_) - | OcrError::MissingField(_) - | OcrError::MissingDocumentUrl - ) - } - Error::Messages(error) => match error { - messages::Error::Auth(source) => auth_is_value_error(source), - _ => error.is_request(), - }, - Error::AudioTranscription(error) => match error { - audio_transcription::Error::Auth(source) => auth_is_value_error(source), - audio_transcription::Error::InvalidProvider(_) - | audio_transcription::Error::InvalidRequest(_) - | audio_transcription::Error::Headers(_) - | audio_transcription::Error::Http(_) - | audio_transcription::Error::InvalidType { .. } - | audio_transcription::Error::MissingField(_) - | audio_transcription::Error::Aws(_) => true, - _ => false, - }, - Error::ChatCompletions(error) => match error { - chat_completions::Error::Auth(source) => auth_is_value_error(source), - chat_completions::Error::InvalidProvider(_) - | chat_completions::Error::InvalidRequest(_) - | chat_completions::Error::Headers(_) - | chat_completions::Error::Http(_) - | chat_completions::Error::InvalidType { .. } - | chat_completions::Error::MissingField(_) - | chat_completions::Error::Aws(_) => true, - _ => false, - }, - Error::Responses(error) => match error { - responses::Error::Auth(source) => auth_is_value_error(source), - responses::Error::InvalidProvider(_) - | responses::Error::InvalidRequest(_) - | responses::Error::Headers(_) => true, - _ => false, - }, - }; - if value_error { - PyValueError::new_err(error.to_string()) +/// A request the caller got wrong is a `ValueError`; anything else is a `RuntimeError`. +pub(crate) fn by_fault(is_request: bool, message: String) -> PyErr { + if is_request { + PyValueError::new_err(message) } else { - PyRuntimeError::new_err(error.to_string()) + PyRuntimeError::new_err(message) } } @@ -96,27 +38,15 @@ pub(crate) fn core_error_to_pyerr(error: Error) -> PyErr { /// Everything raised before the request goes out is safe for the host to retry /// on its own path; anything after it is not, because the provider has already /// done the work and billed for it. -pub(crate) fn chat_completions_error_to_pyerr(error: chat_completions::Error) -> PyErr { - use chat_completions::Error; - match error { - Error::Unsupported(_) - | Error::Auth(_) - | Error::Aws(_) - | Error::InvalidProvider(_) - | Error::InvalidRequest(_) - | Error::InvalidType { .. } - | Error::MissingField(_) - | Error::Headers(_) - | Error::Http(_) - | Error::Transport(TransportError::Connect(_)) => { - RustBridgeDeclined::new_err(error.to_string()) - } - Error::Transport(TransportError::Http { status, body }) => { - RustUpstreamError::new_err((status, body)) - } - Error::Transport(TransportError::Network(message)) | Error::InvalidResponse(message) => { - RustUpstreamError::new_err((0u16, message)) - } +pub(crate) fn chat_completions_error_to_pyerr(error: RouteError) -> PyErr { + match error.phase() { + Phase::BeforeSend => RustBridgeDeclined::new_err(error.to_string()), + Phase::AfterSend => RustUpstreamError::new_err(match error { + RouteError::Transport(TransportError::Http { status, body }) => (status, body), + RouteError::Transport(TransportError::Network(message)) + | RouteError::InvalidResponse(message) => (0u16, message), + other => (0u16, other.to_string()), + }), } } @@ -158,15 +88,14 @@ mod tests { fn missing_api_key_stays_a_runtime_error_while_other_auth_failures_are_value_errors() { Python::initialize(); Python::attach(|py| { - let missing = messages_error_to_pyerr(messages::Error::Auth( - litellm_auth::Error::MissingApiKey { + let missing = + route_error_to_pyerr(RouteError::Auth(litellm_auth::Error::MissingApiKey { provider: "Anthropic", environment_variable: "ANTHROPIC_API_KEY", - }, - )); + })); assert!(missing.is_instance_of::(py)); let invalid = - messages_error_to_pyerr(messages::Error::Auth(litellm_auth::Error::InvalidHeader)); + route_error_to_pyerr(RouteError::Auth(litellm_auth::Error::InvalidHeader)); assert!(invalid.is_instance_of::(py)); }); } diff --git a/litellm-rust/crates/python-bridge/src/http.rs b/litellm-rust/crates/python-bridge/src/http.rs index 3dad3447f45..e74b9d198a0 100644 --- a/litellm-rust/crates/python-bridge/src/http.rs +++ b/litellm-rust/crates/python-bridge/src/http.rs @@ -6,8 +6,8 @@ use std::{ use litellm_core_utils::settings::ProcessEnvironment; use litellm_http::{ - HttpClientConfig, HttpClientPool, HttpSettings, HttpSettingsLayer, Resolution, SslVerify, - TlsSource, Unsupported, + Client, ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, HttpSettingsLayer, + Resolution, SslVerify, TlsSource, Unsupported, media::{PublicDnsResolver, UrlPolicy}, }; use pyo3::{ @@ -80,13 +80,20 @@ fn decode_ssl_verify(field: &Field<'_>) -> Result, ProjectionE Err(field.invalid("a Boolean, Boolean string, CA path, or None")) } -static POOL: LazyLock = - LazyLock::new(|| HttpClientPool::new(Arc::new(PublicDnsResolver))); +static RESOURCES: LazyLock = LazyLock::new(|| { + litellm_core::resources::CoreResources::new(Arc::new(HttpClientPool::new(Arc::new( + PublicDnsResolver, + )))) +}); + +pub(crate) fn resources() -> &'static litellm_core::resources::CoreResources { + &RESOURCES +} static REPORTED_UNSUPPORTED: LazyLock>> = LazyLock::new(Mutex::default); pub(crate) fn pool() -> &'static HttpClientPool { - &POOL + &resources().pool } pub(crate) fn call_config( @@ -97,7 +104,10 @@ pub(crate) fn call_config( let settings = HttpSettings::from_layers([ for_call(call_ssl_verify(kwargs)?, asynchronous), HttpSettingsLayer::from_environment(&ProcessEnvironment), - configured(&PythonSettings::Http.read(py)?)?, + match PythonSettings::Http.read_or_unset(py)? { + Some(snapshot) => configured(&snapshot)?, + None => HttpSettingsLayer::default(), + }, ]) .without_missing_files(&|path: &Path| path.exists()); let resolution = Resolution::from(&settings); @@ -107,6 +117,11 @@ pub(crate) fn call_config( Ok(resolution.config) } +pub(crate) fn host_client(py: Python<'_>, variant: ClientVariant) -> PyResult { + let config = call_config(py, &PyDict::new(py), true)?; + pool().client(&config, variant).map_err(client_error) +} + pub(crate) fn client_error(error: litellm_http::Error) -> PyErr { match error { litellm_http::Error::Read { @@ -143,7 +158,10 @@ fn unreported( } pub(crate) fn url_policy(py: Python<'_>) -> PyResult { - project_url_policy(&PythonSettings::UrlPolicy.read(py)?) + match PythonSettings::UrlPolicy.read_or_unset(py)? { + Some(snapshot) => project_url_policy(&snapshot), + None => Ok(UrlPolicy::default()), + } } fn project_url_policy(snapshot: &Snapshot<'_>) -> PyResult { diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 7c814f540a8..51e112fa1be 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -6,6 +6,7 @@ mod errors; mod http; mod logger; mod marshal; +mod preflight; mod python_settings; mod routes; mod secrets; diff --git a/litellm-rust/crates/callbacks-legacy-python/src/preparation.rs b/litellm-rust/crates/python-bridge/src/preflight.rs similarity index 57% rename from litellm-rust/crates/callbacks-legacy-python/src/preparation.rs rename to litellm-rust/crates/python-bridge/src/preflight.rs index aab654c9893..34813672c09 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/preparation.rs +++ b/litellm-rust/crates/python-bridge/src/preflight.rs @@ -1,9 +1,51 @@ +//! The SDK's request policy the driver runs on every route's keyword view before the host +//! projects from it: credential-name inheritance from `litellm.credential_list`, then the +//! budget and retry-count limits. It is the `@client` prologue after `function_setup` and the +//! deployment hook, and belongs to no callback contract. + use pyo3::{ prelude::*, types::{PyDict, PyList}, }; +use strum::{IntoStaticStr, VariantArray}; -use crate::python::Wrapper; +const MODULE: &str = "litellm.rust_bridge.preflight"; + +/// The litellm globals the preflight still reads through Python. `preflight_contract.json` +/// pins each function's parameters on both sides. +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)] +pub(crate) enum PythonPreflight { + #[strum(serialize = "credential_list")] + CredentialList, + #[strum(serialize = "warn_unknown_credential")] + WarnUnknownCredential, + #[strum(serialize = "check_limits")] + CheckLimits, +} + +impl PythonPreflight { + fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult> + where + A: pyo3::call::PyCallArgs<'py>, + { + py.import(MODULE)?.getattr(<&str>::from(self))?.call1(args) + } +} + +#[cfg(test)] +pub(crate) const PYTHON_CONTRACT: &str = include_str!("../preflight_contract.json"); + +/// Rewrites `arguments` in place, in the order the Python wrapper runs: credentials first, +/// so the limits see the same view the provider request is built from. +pub(crate) fn sdk_preflight(py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<()> { + inherit_credentials(py, arguments, || { + Ok(PythonPreflight::CredentialList + .call(py, ())? + .cast_into::()?) + })?; + PythonPreflight::CheckLimits.call(py, (arguments,))?; + Ok(()) +} struct CredentialEntry<'py>(Bound<'py, PyAny>); @@ -17,22 +59,6 @@ impl<'py> CredentialEntry<'py> { } } -pub fn prepare<'py>( - py: Python<'py>, - kwargs: &Bound<'py, PyDict>, - logger: &crate::PythonLogger, -) -> PyResult> { - let arguments = kwargs.copy()?; - arguments.set_item("litellm_logging_obj", logger.object(py))?; - inherit_credentials(py, &arguments, || { - Ok(Wrapper::CredentialList - .call(py, ())? - .cast_into::()?) - })?; - Wrapper::CheckLimits.call(py, (&arguments,))?; - Ok(arguments) -} - fn inherit_credentials<'py>( py: Python<'py>, arguments: &Bound<'py, PyDict>, @@ -54,7 +80,7 @@ fn inherit_credentials<'py>( .map(|credential| CredentialEntry(credential).name()) .collect::>>()?; let Some(index) = names.iter().position(|name| *name == requested) else { - Wrapper::WarnUnknownCredential.call(py, (requested, names.len()))?; + PythonPreflight::WarnUnknownCredential.call(py, (requested, names.len()))?; return Ok(()); }; let selected = CredentialEntry(credentials.get_item(index)?); @@ -71,7 +97,42 @@ fn inherit_credentials<'py>( #[cfg(test)] mod tests { + use std::collections::BTreeSet; + use std::sync::Mutex; + use super::*; + use strum::VariantArray; + + /// Tests share one interpreter, and the stub module below is global state, so the + /// tests that install it run one at a time. + static PREFLIGHT_MODULE: Mutex<()> = Mutex::new(()); + + /// A fresh stand-in for `litellm.rust_bridge.preflight` that records every call, then + /// `script` run against it with the module bound as `preflight`. + fn preflight_stubs<'py>(py: Python<'py>, script: &std::ffi::CStr) -> Bound<'py, PyDict> { + let locals = PyDict::new(py); + py.run( + c" +import sys +import types + +for name in ('litellm', 'litellm.rust_bridge'): + sys.modules.setdefault(name, types.ModuleType(name)) +preflight = types.ModuleType('litellm.rust_bridge.preflight') +preflight.warnings = [] +preflight.checked = [] +preflight.credential_list = lambda: [] +preflight.warn_unknown_credential = lambda name, loaded: preflight.warnings.append((name, loaded)) +preflight.check_limits = lambda kwargs: preflight.checked.append(kwargs) +sys.modules['litellm.rust_bridge.preflight'] = preflight +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + py.run(script, Some(&locals), Some(&locals)).unwrap(); + locals + } fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> { let locals = PyDict::new(py); @@ -312,4 +373,110 @@ arguments = {'litellm_credential_name': 'ocr-test'} } }); } + + #[test] + fn every_borrowed_function_is_in_the_python_contract() { + Python::initialize(); + Python::attach(|py| { + let contract = litellm_host_python::json_loads(py, PYTHON_CONTRACT.as_bytes()).unwrap(); + let declared: BTreeSet = contract + .bind(py) + .cast::() + .unwrap() + .keys() + .extract() + .map(|names: Vec| names.into_iter().collect()) + .unwrap(); + let called: BTreeSet = PythonPreflight::VARIANTS + .iter() + .map(|&function| <&str>::from(function).to_owned()) + .collect(); + assert_eq!( + called.len(), + PythonPreflight::VARIANTS.len(), + "a function is borrowed twice" + ); + assert_eq!(called, declared); + }); + } + + #[test] + fn an_unknown_name_is_reported_with_the_loaded_count_and_leaves_the_arguments_alone() { + let _guard = PREFLIGHT_MODULE + .lock() + .unwrap_or_else(|error| error.into_inner()); + Python::initialize(); + Python::attach(|py| { + let locals = preflight_stubs( + py, + c" +class Credential: + credential_name = 'listed' + credential_values = {'api_key': 'listed-key'} +preflight.credential_list = lambda: [Credential(), Credential()] +arguments = {'litellm_credential_name': 'missing'} +", + ); + sdk_preflight(py, &argument_dict(&locals)).unwrap(); + py.run( + c" +assert arguments == {'litellm_credential_name': 'missing'}, arguments +assert preflight.warnings == [('missing', 2)], preflight.warnings +assert preflight.checked == [arguments] +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + + #[test] + fn limits_are_checked_on_the_arguments_after_credentials_are_inherited() { + let _guard = PREFLIGHT_MODULE + .lock() + .unwrap_or_else(|error| error.into_inner()); + Python::initialize(); + Python::attach(|py| { + let locals = preflight_stubs( + py, + c" +class Credential: + credential_name = 'ocr-test' + credential_values = {'api_key': 'inherited'} +preflight.credential_list = lambda: [Credential()] +rejection = RuntimeError('Max retries per request hit!') +def check_limits(arguments): + preflight.checked.append(dict(arguments)) + raise rejection +preflight.check_limits = check_limits +arguments = {'litellm_credential_name': 'ocr-test'} +", + ); + let error = sdk_preflight(py, &argument_dict(&locals)).unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("rejection").unwrap().unwrap()) + ); + py.run( + c" +assert preflight.checked == [{'litellm_credential_name': 'ocr-test', 'api_key': 'inherited'}], preflight.checked +assert arguments['api_key'] == 'inherited' +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + + fn argument_dict<'py>(locals: &Bound<'py, PyDict>) -> Bound<'py, PyDict> { + locals + .get_item("arguments") + .unwrap() + .unwrap() + .cast_into::() + .unwrap() + } } diff --git a/litellm-rust/crates/python-bridge/src/python_settings.rs b/litellm-rust/crates/python-bridge/src/python_settings.rs index abf664b795d..8f49444f730 100644 --- a/litellm-rust/crates/python-bridge/src/python_settings.rs +++ b/litellm-rust/crates/python-bridge/src/python_settings.rs @@ -1,11 +1,14 @@ -use pyo3::prelude::*; +use pyo3::{exceptions::PyModuleNotFoundError, prelude::*}; +use strum::IntoStaticStr; use crate::coercion::{FieldSpec, ProjectionError}; const MODULE: &str = "litellm.rust_bridge.settings"; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq)] +#[strum(serialize_all = "snake_case")] pub(crate) enum PythonSettings { + #[strum(serialize = "http_settings")] Http, UrlPolicy, ProviderDefaults, @@ -26,13 +29,7 @@ impl Snapshot<'_> { impl PythonSettings { pub(crate) fn name(self) -> &'static str { - match self { - Self::Http => "http_settings", - Self::UrlPolicy => "url_policy", - Self::ProviderDefaults => "provider_defaults", - Self::SecretManager => "secret_manager", - Self::SecretManagerBinding => "secret_manager_binding", - } + self.into() } pub(crate) fn read(self, py: Python<'_>) -> PyResult> { @@ -40,15 +37,45 @@ impl PythonSettings { Ok(Snapshot { group: self, value }) } + /// Reads the accessor, or `None` when the litellm package is not installed + /// (a bare extension module), meaning there are no configured values. + pub(crate) fn read_or_unset(self, py: Python<'_>) -> PyResult>> { + match self.read(py) { + Ok(snapshot) => Ok(Some(snapshot)), + Err(error) => { + if missing_module(py, &error, "litellm")? { + Ok(None) + } else { + Err(error) + } + } + } + } + #[cfg(test)] pub(crate) fn snapshot(self, value: Bound<'_, PyAny>) -> Snapshot<'_> { Snapshot { group: self, value } } } +fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult { + if !error.is_instance_of::(py) { + return Ok(false); + } + Ok(error + .value(py) + .getattr("name")? + .extract::>()? + .is_some_and(|name| name == expected)) +} + #[cfg(test)] mod tests { - use pyo3::{exceptions::PyRuntimeError, prelude::*, types::PyDict}; + use pyo3::{ + exceptions::{PyImportError, PyModuleNotFoundError, PyRuntimeError}, + prelude::*, + types::PyDict, + }; use super::PythonSettings; use crate::coercion::FieldSpec; @@ -140,4 +167,153 @@ values = (Descriptor(), SimpleNamespace(flag=Truth())) ); }); } + + #[test] + fn read_or_unset_returns_none_when_litellm_is_missing() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + c" +import sys +class MissingLitellm: + def find_spec(self, fullname, path=None, target=None): + if fullname == 'litellm': + raise ModuleNotFoundError('No module named litellm', name='litellm') +finder = MissingLitellm() +previous_litellm = sys.modules.get('litellm') +had_litellm = 'litellm' in sys.modules +sys.meta_path.insert(0, finder) +sys.modules.pop('litellm', None) +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + let result = PythonSettings::Http.read_or_unset(py); + assert!(result.unwrap().is_none()); + py.run( + c" +sys.meta_path.remove(finder) +if had_litellm: + sys.modules['litellm'] = previous_litellm +else: + sys.modules.pop('litellm', None) +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + + #[test] + fn read_or_unset_propagates_nested_module_not_found_errors() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + c" +import sys +import types +previous_modules = { + name: sys.modules[name] + for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.settings') + if name in sys.modules +} +litellm = types.ModuleType('litellm') +litellm.__path__ = [] +rust_bridge = types.ModuleType('litellm.rust_bridge') +rust_bridge.__path__ = [] +settings = types.ModuleType('litellm.rust_bridge.settings') +def http_settings(): + raise ModuleNotFoundError('No module named certifi', name='certifi') +settings.http_settings = http_settings +litellm.rust_bridge = rust_bridge +rust_bridge.settings = settings +sys.modules['litellm'] = litellm +sys.modules['litellm.rust_bridge'] = rust_bridge +sys.modules['litellm.rust_bridge.settings'] = settings +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + let error = match PythonSettings::Http.read_or_unset(py) { + Ok(_) => panic!("nested module errors must propagate"), + Err(error) => error, + }; + assert!(error.is_instance_of::(py)); + assert_eq!( + error + .value(py) + .getattr("name") + .unwrap() + .extract::() + .unwrap(), + "certifi" + ); + py.run( + c" +for name in ('litellm.rust_bridge.settings', 'litellm.rust_bridge', 'litellm'): + sys.modules.pop(name, None) +sys.modules.update(previous_modules) +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + + #[test] + fn read_or_unset_propagates_import_errors() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + c" +import sys +import types +previous_modules = { + name: sys.modules[name] + for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.settings') + if name in sys.modules +} +litellm = types.ModuleType('litellm') +litellm.__path__ = [] +rust_bridge = types.ModuleType('litellm.rust_bridge') +rust_bridge.__path__ = [] +settings = types.ModuleType('litellm.rust_bridge.settings') +def http_settings(): + raise ImportError('cannot import name setting') +settings.http_settings = http_settings +litellm.rust_bridge = rust_bridge +rust_bridge.settings = settings +sys.modules['litellm'] = litellm +sys.modules['litellm.rust_bridge'] = rust_bridge +sys.modules['litellm.rust_bridge.settings'] = settings +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + let error = match PythonSettings::Http.read_or_unset(py) { + Ok(_) => panic!("import errors must propagate"), + Err(error) => error, + }; + assert!(error.is_instance_of::(py)); + assert_eq!(error.to_string(), "ImportError: cannot import name setting"); + py.run( + c" +for name in ('litellm.rust_bridge.settings', 'litellm.rust_bridge', 'litellm'): + sys.modules.pop(name, None) +sys.modules.update(previous_modules) +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } } diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index dec4dcea21c..ad80659de92 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -3,15 +3,17 @@ use litellm_core::audio_transcription::{ Error, audio_transcription as run_audio_transcription, types::AudioTranscriptionRequest, }; use litellm_host_python::from_py_argument; -use pyo3::prelude::*; +use litellm_http::HttpClientConfig; +use pyo3::{prelude::*, types::PyDict}; use serde_json::{Map, Value}; use crate::{ - errors::audio_transcription_error_to_pyerr, + errors::route_error_to_pyerr, marshal::{RouteOptions, extra_headers_argument, optional_params_argument, optional_timeout}, }; async fn execute( + config: HttpClientConfig, audio: Value, optional_params: Map, options: RouteOptions, @@ -24,16 +26,20 @@ async fn execute( extra_headers, timeout, } = options; - run_audio_transcription(AudioTranscriptionRequest { - model: &model, - audio, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - optional_params, - timeout, - }) + run_audio_transcription( + crate::http::resources(), + &config, + AudioTranscriptionRequest { + model: &model, + audio, + api_key: api_key.as_deref(), + api_base: api_base.as_deref(), + custom_llm_provider: custom_llm_provider.as_deref(), + extra_headers, + optional_params, + timeout, + }, + ) .await } @@ -62,10 +68,11 @@ pub(crate) fn transcription( extra_headers, timeout: optional_timeout(timeout_seconds), }; + let config = crate::http::call_config(py, &PyDict::new(py), false)?; run_sync( py, - execute(audio, optional_params.unwrap_or_default(), options), - audio_transcription_error_to_pyerr, + execute(config, audio, optional_params.unwrap_or_default(), options), + route_error_to_pyerr, ) } @@ -94,9 +101,10 @@ pub(crate) fn atranscription<'py>( extra_headers, timeout: optional_timeout(timeout_seconds), }; + let config = crate::http::call_config(py, &PyDict::new(py), true)?; run_async( py, - execute(audio, optional_params.unwrap_or_default(), options), - audio_transcription_error_to_pyerr, + execute(config, audio, optional_params.unwrap_or_default(), options), + route_error_to_pyerr, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index b96b12bfc43..f4fd53c61b0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -7,6 +7,7 @@ use litellm_core::chat_completions::{ types::ChatCompletionsRequest, }; use litellm_host_python::from_py_argument; +use litellm_http::HttpClientConfig; use litellm_types::utils::ChatCompletionsResponse; use pyo3::prelude::*; use serde_json::{Map, Value}; @@ -20,6 +21,7 @@ use crate::{ }; async fn execute( + config: HttpClientConfig, messages: Vec, optional_params: Map, options: RouteOptions, @@ -32,16 +34,20 @@ async fn execute( extra_headers, timeout, } = options; - run_chat_completions(ChatCompletionsRequest { - model: &model, - messages: Value::Array(messages), - optional_params, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }) + run_chat_completions( + crate::http::resources(), + &config, + ChatCompletionsRequest { + model: &model, + messages: Value::Array(messages), + optional_params, + api_key: api_key.as_deref(), + api_base: api_base.as_deref(), + custom_llm_provider: custom_llm_provider.as_deref(), + extra_headers, + timeout, + }, + ) .await } @@ -87,9 +93,15 @@ pub(crate) fn chat_completions( extra_headers, timeout: optional_timeout(timeout_seconds), }; + let config = crate::http::call_config(py, &PyDict::new(py), false)?; run_sync( py, - execute(messages, optional_params.unwrap_or_default(), options), + execute( + config, + messages, + optional_params.unwrap_or_default(), + options, + ), chat_completions_error_to_pyerr, ) } @@ -119,9 +131,15 @@ pub(crate) fn achat_completions<'py>( extra_headers, timeout: optional_timeout(timeout_seconds), }; + let config = crate::http::call_config(py, &PyDict::new(py), true)?; run_async( py, - execute(messages, optional_params.unwrap_or_default(), options), + execute( + config, + messages, + optional_params.unwrap_or_default(), + options, + ), chat_completions_error_to_pyerr, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index a253f4f5670..cf634ab3fc3 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -2,9 +2,8 @@ use std::convert::Infallible; use bytes::Bytes; use litellm_core::messages::{ - Error, - route::{Messages, MessagesCall, MessagesOutput, MessagesStreamHead}, - types::MessagesShaping, + Error, MessagesCall, MessagesShaping, messages_body, + route::{Messages, MessagesOutput, MessagesStreamHead}, }; use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py}; use litellm_http::transport::Error as TransportError; @@ -18,7 +17,7 @@ use pyo3::{ use serde_json::{Map, Value}; use crate::{ - errors::{RustUpstreamError, messages_error_to_pyerr}, + errors::{RustUpstreamError, route_error_to_pyerr}, marshal::{optional_timeout, python_timeout_seconds}, }; @@ -76,7 +75,12 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult { error.value(py).setattr(REQUEST_ERROR_MARKER, true)?; Ok(error) } - other => Ok(messages_error_to_pyerr(other)), + Error::MissingField(field) => { + let error = PyValueError::new_err(format!("missing required field: {field}")); + error.value(py).setattr(REQUEST_ERROR_MARKER, true)?; + Ok(error) + } + other => Ok(route_error_to_pyerr(other)), } } @@ -91,7 +95,11 @@ impl MessagesPythonHost { Self { request } } - fn projection(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult { + fn projection( + &self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> PyResult> { let request = self.request.bind(py); let argument = |name: &str| -> PyResult>> { Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none())) @@ -123,17 +131,20 @@ impl MessagesPythonHost { .flatten(); let custom_llm_provider = string("custom_llm_provider")?; let shaping = self.shaping(py, &model, custom_llm_provider.as_deref(), arguments)?; - Ok(MessagesCall { - model, + let api_key = string("api_key")?; + let api_base = string("api_base")?; + let extra_headers = self.merged_headers(py, arguments)?; + let provider_specific_header = self.provider_specific_header(py, arguments)?; + Ok(messages_body(body).map(|body| MessagesCall { body, - api_key: string("api_key")?, - api_base: string("api_base")?, - extra_headers: self.merged_headers(py, arguments)?, - provider_specific_header: self.provider_specific_header(py, arguments)?, + api_key, + api_base, + extra_headers, + provider_specific_header, custom_llm_provider, timeout: optional_timeout(timeout), shaping, - }) + })) } fn merged_headers( @@ -220,7 +231,8 @@ impl ProtocolHost for MessagesPythonHost { arguments: &Bound<'_, PyDict>, ) -> Result> { self.projection(py, arguments) - .map_err(|error| InvokeError::Python(self.map_failure(py, error))) + .map_err(|error| InvokeError::Python(self.map_failure(py, error)))? + .map_err(InvokeError::Native) } fn invoke(&mut self, _: Python<'_>, op: Infallible) -> Result<(), InvokeError> { @@ -306,6 +318,7 @@ mod tests { #[rstest] #[case::rejected_request(Error::InvalidRequest("does not support top_k=5".into()), true)] + #[case::missing_field(Error::MissingField("max_tokens"), true)] #[case::unresolvable_provider(Error::InvalidProvider("openai".into()), false)] #[case::upstream_failure( Error::Transport(TransportError::Http { status: 400, body: "bad".into() }), diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index dae8623979a..96ca9eebecb 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -27,12 +27,16 @@ fn run_messages( asynchronous: bool, ) -> PyResult> { let secrets = crate::secrets::source(py)?; + let config = crate::http::call_config(py, &kwargs, asynchronous)?; + let machine = messages_machine(crate::http::resources(), &config, secrets) + .map_err(crate::http::client_error)?; run_legacy_call( py, SURFACE, PublicCall::capture(&request, &args, &kwargs)?, - crate::logger::LoggedMachine::new(messages_machine(secrets)), + crate::logger::LoggedMachine::new(machine), MessagesPythonHost::new(request.unbind()), + crate::preflight::sdk_preflight, asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs index b0a6acdebfd..2068b6a6e4b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs @@ -4,7 +4,7 @@ use pyo3::{ prelude::*, }; -use crate::errors::{RustUpstreamError, core_error_to_pyerr}; +use crate::errors::{RustUpstreamError, by_fault}; pub(super) fn to_pyerr(error: Error) -> PyErr { let status = error.http_status_code(); @@ -19,7 +19,7 @@ pub(super) fn to_pyerr(error: Error) -> PyErr { upstream_error(py, status, body, Vec::new())? } Error::RequestFormat => { - let error = core_error_to_pyerr(Error::RequestFormat.into()); + let error = by_fault(true, Error::RequestFormat.to_string()); error .value(py) .setattr("ocr_request_format_error", true) @@ -30,13 +30,25 @@ pub(super) fn to_pyerr(error: Error) -> PyErr { PyFileNotFoundError::new_err(format!("File not found: {}", path.display())) } Error::FileRead { source, .. } => PyOSError::new_err(source.to_string()), - other => core_error_to_pyerr(other.into()), + other => by_fault(is_request(&other), other.to_string()), }) }) .unwrap_or_else(|error| error); attach_status(mapped, status) } +fn is_request(error: &Error) -> bool { + error.is_request() + || matches!( + error, + Error::Auth(_) + | Error::InvalidProvider(_) + | Error::InvalidRequest(_) + | Error::MissingField(_) + | Error::MissingDocumentUrl + ) +} + fn upstream_error( py: Python<'_>, status: u16, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index a4f2bf851d7..d7c54e996ff 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -3,15 +3,12 @@ mod errors; mod host; mod project; -use std::sync::LazyLock; - use host::OcrPythonHost; -use litellm_auth_gcp::VertexAuth; use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; use litellm_core::ocr::{provider_config, route::ocr_machine}; use litellm_core_utils::settings::ProcessEnvironment; use litellm_host_python::to_py; -use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings}; +use litellm_llms::base_llm::ocr::settings::OcrSettings; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, @@ -44,8 +41,6 @@ const ASYNC_SURFACE: LegacySurface = LegacySurface { ..SURFACE }; -static VERTEX_AUTH: LazyLock = LazyLock::new(VertexAuth::default); - fn run_ocr( py: Python<'_>, request: Bound<'_, PyAny>, @@ -55,21 +50,16 @@ fn run_ocr( ) -> PyResult> { let secrets = secrets::source(py)?; let config = http::call_config(py, &kwargs, asynchronous)?; - let client = OcrClient::new( - http::pool(), - &config, - http::url_policy(py)?, - VERTEX_AUTH.clone(), - ocr_settings(py)?, - secrets, - ) - .map_err(http::client_error)?; + let client = http::resources() + .ocr_client(&config, http::url_policy(py)?, ocr_settings(py)?, secrets) + .map_err(http::client_error)?; run_legacy_call( py, if asynchronous { ASYNC_SURFACE } else { SURFACE }, PublicCall::capture(&request, &args, &kwargs)?, crate::logger::LoggedMachine::new(ocr_machine(client)), OcrPythonHost::new(request.unbind()), + crate::preflight::sdk_preflight, asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses.rs index 5995d64649b..bf17ef6edde 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses.rs @@ -6,7 +6,7 @@ use pyo3::{ use serde_json::Value; use crate::{ - errors::{RustBridgeDeclined, responses_error_to_pyerr}, + errors::{RustBridgeDeclined, route_error_to_pyerr}, marshal::{marshal_headers, optional_timeout}, }; @@ -57,7 +57,7 @@ impl ResponsesWebSocketConnection { crate::logger::run_async_value(py, async move { let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout) .await - .map_err(responses_error_to_pyerr)?; + .map_err(route_error_to_pyerr)?; Ok(ResponsesWebSocketConnection { inner }) }) } @@ -65,24 +65,21 @@ impl ResponsesWebSocketConnection { fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult> { let inner = self.inner.clone(); crate::logger::run_async_value(py, async move { - inner - .send_text(text) - .await - .map_err(responses_error_to_pyerr) + inner.send_text(text).await.map_err(route_error_to_pyerr) }) } fn recv_text<'py>(&self, py: Python<'py>) -> PyResult> { let inner = self.inner.clone(); crate::logger::run_async_value(py, async move { - inner.recv_text().await.map_err(responses_error_to_pyerr) + inner.recv_text().await.map_err(route_error_to_pyerr) }) } fn close<'py>(&self, py: Python<'py>) -> PyResult> { let inner = self.inner.clone(); crate::logger::run_async_value(py, async move { - inner.close().await.map_err(responses_error_to_pyerr) + inner.close().await.map_err(route_error_to_pyerr) }) } } diff --git a/litellm-rust/crates/python-bridge/src/secrets/callback.rs b/litellm-rust/crates/python-bridge/src/secrets/callback.rs index 82bb4443f98..6b61ca22dcd 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/callback.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/callback.rs @@ -73,17 +73,7 @@ impl PythonClient { /// The `KeyManagementSystem` value as Python spells it. fn python_name(system: KeyManagementSystem) -> &'static str { - match system { - KeyManagementSystem::GoogleKms => "google_kms", - KeyManagementSystem::AzureKeyVault => "azure_key_vault", - KeyManagementSystem::AwsSecretManager => "aws_secret_manager", - KeyManagementSystem::GoogleSecretManager => "google_secret_manager", - KeyManagementSystem::HashicorpVault => "hashicorp_vault", - KeyManagementSystem::Cyberark => "cyberark", - KeyManagementSystem::Local => "local", - KeyManagementSystem::AwsKms => "aws_kms", - KeyManagementSystem::Custom => "custom", - } + system.into() } impl ExternalSecretManager for PythonSecretManager { @@ -210,7 +200,7 @@ handler.get_secret_from_manager = get_secret_from_manager KeyManagementSettings::default(), )), Arc::new(move |_: &str| fallback.map(str::to_owned)), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ) .with_failure_policy(FailurePolicy::EnvironmentFallback); (resolver, locals, handler) diff --git a/litellm-rust/crates/python-bridge/src/secrets/mod.rs b/litellm-rust/crates/python-bridge/src/secrets/mod.rs index 439f9ddddd1..ed54c306397 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/mod.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/mod.rs @@ -27,7 +27,12 @@ const NATIVE: FieldSpec = FieldSpec::new("native", |field| field.schema_bo pub(crate) fn source(py: Python<'_>) -> PyResult> { if PythonSettings::SecretManager.read(py)?.read(&NATIVE)? { let context = litellm_host_python::PythonContext::capture(py)?; - return Ok(Arc::new(ResolvedSecrets::new(config::read(py)?, context))); + let client = crate::http::host_client(py, litellm_http::ClientVariant::Provider)?; + return Ok(Arc::new(ResolvedSecrets::new( + config::read(py)?, + context, + client, + ))); } Ok(Arc::new(PythonSecrets::new(py)?)) } diff --git a/litellm-rust/crates/python-bridge/src/secrets/resolved.rs b/litellm-rust/crates/python-bridge/src/secrets/resolved.rs index 40b187c99de..5a606ab1039 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/resolved.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/resolved.rs @@ -3,6 +3,7 @@ use std::sync::Arc; use futures_util::future::BoxFuture; use litellm_core_utils::settings::ProcessEnvironment; use litellm_host_python::PythonContext; +use litellm_http::Client; use litellm_secrets::source::SecretSource; use litellm_secrets::{ Error, FailurePolicy, OidcResolver, SecretManagerState, SecretResolver, SecretValue, @@ -15,16 +16,20 @@ pub(crate) struct ResolvedSecrets { } impl ResolvedSecrets { - pub(crate) fn new(snapshot: SecretManagerSnapshot, context: PythonContext) -> Self { - Self::from_state(snapshot.into_state(context)) + pub(crate) fn new( + snapshot: SecretManagerSnapshot, + context: PythonContext, + client: Client, + ) -> Self { + Self::from_state(snapshot.into_state(context), client) } - fn from_state(state: Arc) -> Self { + fn from_state(state: Arc, client: Client) -> Self { Self { resolver: SecretResolver::new_python_compatible( state, Arc::new(ProcessEnvironment), - OidcResolver::default(), + OidcResolver::new(client), ) .with_failure_policy(FailurePolicy::EnvironmentFallback), } @@ -79,7 +84,7 @@ mod tests { } async fn resolve(state: Arc, name: &'static str) -> Option { - ResolvedSecrets::from_state(state) + ResolvedSecrets::from_state(state, litellm_http::Client::plain_for_test()) .resolve(&[name]) .await .unwrap() @@ -175,7 +180,10 @@ mod tests { .expect(1) .mount(&server) .await; - let source = ResolvedSecrets::from_state(state(&server, KeyManagementSettings::default())); + let source = ResolvedSecrets::from_state( + state(&server, KeyManagementSettings::default()), + litellm_http::Client::plain_for_test(), + ); let snapshot = source.resolve(&[declared]).await.unwrap(); assert_eq!(snapshot.get(undeclared), None); let result = source @@ -238,9 +246,12 @@ mod tests { #[tokio::test] async fn oidc_failures_are_not_converted_to_missing_secrets() { - let result = ResolvedSecrets::from_state(Arc::new(SecretManagerState::default())) - .resolve(&["oidc/"]) - .await; + let result = ResolvedSecrets::from_state( + Arc::new(SecretManagerState::default()), + litellm_http::Client::plain_for_test(), + ) + .resolve(&["oidc/"]) + .await; assert!(matches!(result, Err(litellm_secrets::Error::InvalidOidc))); } diff --git a/litellm-rust/crates/python-bridge/src/secrets/runtime.rs b/litellm-rust/crates/python-bridge/src/secrets/runtime.rs index 1a89130ee82..4d2e88115c8 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/runtime.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/runtime.rs @@ -10,6 +10,7 @@ use litellm_secrets_types::PythonSecretRead; use pyo3::{ exceptions::{PyAttributeError, PyRuntimeError, PyValueError}, prelude::*, + types::PyDict, }; #[derive(Clone, PartialEq)] @@ -44,10 +45,18 @@ impl NativeSecretManager { let system = configuration.system; let settings = configuration.settings.clone(); let enterprise_enabled = configuration.enterprise_enabled; + let http_config = crate::http::call_config(py, &PyDict::new(py), false)?; let backend = run_sync_value(py, async move { - load_native_manager(system, settings, environment, enterprise_enabled) - .await - .map_err(|error| PyValueError::new_err(error.to_string())) + load_native_manager( + crate::http::pool(), + &http_config, + system, + settings, + environment, + enterprise_enabled, + ) + .await + .map_err(|error| PyValueError::new_err(error.to_string())) })?; Ok(Self { backend, diff --git a/litellm-rust/crates/router/Cargo.toml b/litellm-rust/crates/router/Cargo.toml new file mode 100644 index 00000000000..cc6972a066e --- /dev/null +++ b/litellm-rust/crates/router/Cargo.toml @@ -0,0 +1,13 @@ +[package] +name = "litellm-router" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-config.workspace = true +litellm-core.workspace = true + +[dev-dependencies] +rstest.workspace = true diff --git a/litellm-rust/crates/router/README.md b/litellm-rust/crates/router/README.md new file mode 100644 index 00000000000..ed5b2966fe8 --- /dev/null +++ b/litellm-rust/crates/router/README.md @@ -0,0 +1,5 @@ +`litellm-router` scaffolds the model-list setup and deployment lookup portion of Python's `litellm.Router`. `Router::from_model_list(&config.model_list)` maps configured public names to provider deployments. Programmatic callers can collect `(String, Deployment)` entries into a `Router` + +Lookup is exact and returns `None` for an unknown name. This extraction preserves the gateway's existing behavior: the last entry wins when public names repeat. Multiple deployments per model group, routing strategies, retries, cooldowns, and fallbacks are not implemented yet + +The router owns deployment configuration and selection. The gateway handles HTTP errors and responses, while `core` executes provider calls and resolves credentials diff --git a/litellm-rust/crates/router/src/deployment.rs b/litellm-rust/crates/router/src/deployment.rs new file mode 100644 index 00000000000..9234728db28 --- /dev/null +++ b/litellm-rust/crates/router/src/deployment.rs @@ -0,0 +1,13 @@ +use std::time::Duration; + +use litellm_core::messages::MessagesShaping; + +#[derive(Clone, Debug, Default)] +pub struct Deployment { + pub model: String, + pub api_key: Option, + pub api_base: Option, + pub custom_llm_provider: Option, + pub timeout: Option, + pub shaping: MessagesShaping, +} diff --git a/litellm-rust/crates/router/src/lib.rs b/litellm-rust/crates/router/src/lib.rs new file mode 100644 index 00000000000..da33bfb04bd --- /dev/null +++ b/litellm-rust/crates/router/src/lib.rs @@ -0,0 +1,44 @@ +mod deployment; + +use std::collections::HashMap; + +use litellm_config::Model; + +pub use deployment::Deployment; + +#[derive(Clone, Debug, Default)] +pub struct Router(HashMap); + +impl Router { + pub fn from_model_list(model_list: &[Model]) -> Self { + model_list + .iter() + .map(|model| { + ( + model.model_name.clone(), + Deployment { + model: model.litellm_params.model.clone(), + api_key: model + .litellm_params + .api_key + .as_ref() + .map(|value| value.expose().to_string()), + api_base: model.litellm_params.api_base.clone(), + custom_llm_provider: model.litellm_params.custom_llm_provider.clone(), + ..Deployment::default() + }, + ) + }) + .collect() + } + + pub fn get(&self, model_name: &str) -> Option<&Deployment> { + self.0.get(model_name) + } +} + +impl FromIterator<(String, Deployment)> for Router { + fn from_iter>(entries: I) -> Self { + Self(entries.into_iter().collect()) + } +} diff --git a/litellm-rust/crates/router/tests/router.rs b/litellm-rust/crates/router/tests/router.rs new file mode 100644 index 00000000000..2d16f102de3 --- /dev/null +++ b/litellm-rust/crates/router/tests/router.rs @@ -0,0 +1,94 @@ +use std::time::Duration; + +use litellm_config::Config; +use litellm_core::messages::MessagesShaping; +use litellm_router::{Deployment, Router}; +use rstest::rstest; + +#[rstest] +#[case::minimal("")] +#[case::configured( + "api_key: test-key\n api_base: https://provider.example/v1\n custom_llm_provider: test-provider" +)] +#[case::secret_reference("api_key: os.environ/ROUTER_TEST_API_KEY")] +fn configuration_preserves_deployment_parameters(#[case] parameters: &str) { + let config = Config::from_yaml(&format!( + "model_list:\n - model_name: public-model\n litellm_params:\n model: provider/model\n {parameters}" + )) + .unwrap(); + let router = Router::from_model_list(&config.model_list); + let deployment = router.get(&config.model_list[0].model_name).unwrap(); + let params = &config.model_list[0].litellm_params; + + assert_eq!(deployment.model, params.model); + assert_eq!( + deployment.api_key.as_deref(), + params.api_key.as_ref().map(|key| key.expose()) + ); + assert_eq!(deployment.api_base, params.api_base); + assert_eq!(deployment.custom_llm_provider, params.custom_llm_provider); + assert_eq!(deployment.timeout, Deployment::default().timeout); + assert_eq!(deployment.shaping, Deployment::default().shaping); +} + +#[rstest] +#[case::first("public-a", Some("provider/a"))] +#[case::second("public-b", Some("provider/b"))] +#[case::unknown("missing", None)] +#[case::provider_name_is_not_an_alias("provider/a", None)] +#[case::case_sensitive("PUBLIC-A", None)] +fn lookup_uses_public_names(#[case] name: &str, #[case] expected: Option<&str>) { + let config = Config::from_yaml( + "model_list: + - model_name: public-a + litellm_params: + model: provider/a + - model_name: public-b + litellm_params: + model: provider/b", + ) + .unwrap(); + let router = Router::from_model_list(&config.model_list); + + assert_eq!(router.get(name).map(|entry| entry.model.as_str()), expected); +} + +#[rstest] +fn empty_configuration_has_no_deployment() { + let config = Config::from_yaml("model_list: []").unwrap(); + + assert!( + Router::from_model_list(&config.model_list) + .get("") + .is_none() + ); + assert!(Router::default().get("unknown").is_none()); +} + +#[rstest] +fn programmatic_deployments_preserve_overrides_and_last_entry_wins() { + let deployment = Deployment { + model: "provider/selected".into(), + api_key: Some("test-key".into()), + api_base: Some("https://provider.example/v1".into()), + custom_llm_provider: Some("test-provider".into()), + timeout: Some(Duration::from_secs(7)), + shaping: MessagesShaping { + drop_params: true, + additional_drop_params: vec!["metadata.test".into()], + ..Default::default() + }, + }; + let router = Router::from_iter([ + ("public-model".into(), Deployment::default()), + ("public-model".into(), deployment.clone()), + ]); + let selected = router.get("public-model").unwrap(); + + assert_eq!(selected.model, deployment.model); + assert_eq!(selected.api_key, deployment.api_key); + assert_eq!(selected.api_base, deployment.api_base); + assert_eq!(selected.custom_llm_provider, deployment.custom_llm_provider); + assert_eq!(selected.timeout, deployment.timeout); + assert_eq!(selected.shaping, deployment.shaping); +} diff --git a/litellm-rust/crates/secrets-aws/src/auth.rs b/litellm-rust/crates/secrets-aws/src/auth.rs index 0c32eb00989..06019cc3c9d 100644 --- a/litellm-rust/crates/secrets-aws/src/auth.rs +++ b/litellm-rust/crates/secrets-aws/src/auth.rs @@ -2,9 +2,8 @@ use std::sync::Arc; use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future}; use litellm_auth_aws::{ - AwsAuthConfig, + AwsAuthConfig, AwsAuthService, constants::{AWS_DEFAULT_REGION, AWS_REGION, AWS_REGION_NAME}, - resolve_credentials, }; use litellm_core_utils::settings::Lookup; use litellm_secrets_types::{AwsOperationContext, KeyManagementSettings}; @@ -13,6 +12,7 @@ use crate::Error; #[derive(Clone)] pub(crate) struct Credentials { + auth: AwsAuthService, config: AwsAuthConfig, environment: Arc, } @@ -22,15 +22,22 @@ impl Credentials { settings: &KeyManagementSettings, environment: Arc, ) -> Self { - Self::with_context(settings, environment, &AwsOperationContext::default()) + Self::with_context( + AwsAuthService::default(), + settings, + environment, + &AwsOperationContext::default(), + ) } pub(crate) fn with_context( + auth: AwsAuthService, settings: &KeyManagementSettings, environment: Arc, context: &AwsOperationContext, ) -> Self { Self { + auth, config: AwsAuthConfig { access_key_id: context .access_key_id @@ -69,7 +76,8 @@ impl ProvideCredentials for Credentials { Self: 'a, { future::ProvideCredentials::new(async { - resolve_credentials(self.config.clone(), &|name| self.environment.get(name)) + self.auth + .resolve_credentials(self.config.clone(), &|name| self.environment.get(name)) .await .map_err(|_| { CredentialsError::provider_error("secret manager authentication failed") diff --git a/litellm-rust/crates/secrets-aws/src/secret_manager.rs b/litellm-rust/crates/secrets-aws/src/secret_manager.rs index 508d7da15c1..2220e387838 100644 --- a/litellm-rust/crates/secrets-aws/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-aws/src/secret_manager.rs @@ -39,6 +39,7 @@ pub struct AwsSecretsManagerV2 { #[derive(Clone)] struct ContextClientFactory { + auth: litellm_auth_aws::AwsAuthService, settings: KeyManagementSettings, environment: Arc, endpoint_url: Option, diff --git a/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs b/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs index aac998c65ab..0f0032ed64f 100644 --- a/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs +++ b/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs @@ -22,6 +22,7 @@ impl AwsSecretsManagerV2 { return Ok(None); } let context_client_factory = ContextClientFactory { + auth: litellm_auth_aws::AwsAuthService::default(), settings: settings.clone(), environment: environment.clone(), endpoint_url: environment @@ -90,6 +91,7 @@ impl ContextClientFactory { self.environment.as_ref(), )?)) .credentials_provider(auth::Credentials::with_context( + self.auth.clone(), &settings, self.environment.clone(), context, diff --git a/litellm-rust/crates/secrets-azure/Cargo.toml b/litellm-rust/crates/secrets-azure/Cargo.toml index 7e8a79f89ef..efdf681e2bc 100644 --- a/litellm-rust/crates/secrets-azure/Cargo.toml +++ b/litellm-rust/crates/secrets-azure/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true tokio.workspace = true litellm-auth-azure.workspace = true litellm-auth-types.workspace = true @@ -18,6 +19,7 @@ veil.workspace = true percent-encoding = "2.3" [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } wiremock = "0.6.5" rstest.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/secrets-azure/src/key_vault.rs b/litellm-rust/crates/secrets-azure/src/key_vault.rs index 13e59f8e4ac..095c451927c 100644 --- a/litellm-rust/crates/secrets-azure/src/key_vault.rs +++ b/litellm-rust/crates/secrets-azure/src/key_vault.rs @@ -19,7 +19,7 @@ const PATH_SEGMENT: &AsciiSet = &NON_ALPHANUMERIC #[derive(Clone)] pub struct AzureKeyVault { - client: reqwest::Client, + client: litellm_http::Client, vault: reqwest::Url, auth: Arc, inputs: Arc, @@ -33,7 +33,7 @@ struct SecretResponse { impl AzureKeyVault { pub fn with_client( - client: reqwest::Client, + client: litellm_http::Client, vault: reqwest::Url, environment: Arc, ) -> Result { @@ -57,7 +57,10 @@ impl AzureKeyVault { }) } - pub fn new(environment: Arc) -> Result { + pub fn new( + client: litellm_http::Client, + environment: Arc, + ) -> Result { let value = environment .get(AZURE_KEY_VAULT_URI) .ok_or(Error::MissingEnvironment(AZURE_KEY_VAULT_URI))?; @@ -65,7 +68,7 @@ impl AzureKeyVault { if vault.scheme() != "https" || vault.host_str().is_none() { return Err(Error::VaultUri); } - Self::with_client(reqwest::Client::new(), vault, environment) + Self::with_client(client, vault, environment) } pub fn scope(&self) -> &str { diff --git a/litellm-rust/crates/secrets-azure/tests/key_vault.rs b/litellm-rust/crates/secrets-azure/tests/key_vault.rs index a21149db345..fcbc46092e1 100644 --- a/litellm-rust/crates/secrets-azure/tests/key_vault.rs +++ b/litellm-rust/crates/secrets-azure/tests/key_vault.rs @@ -11,7 +11,7 @@ use wiremock::{ fn manager(server: &MockServer) -> AzureKeyVault { AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "fake".to_owned())), ) @@ -130,11 +130,14 @@ fn new_validates_vault_environment( #[case] uri: Option<&'static str>, #[case] missing_environment: bool, ) { - let result = AzureKeyVault::new(Arc::new(move |name: &str| { - (name == "AZURE_KEY_VAULT_URI") - .then(|| uri.map(str::to_owned)) - .flatten() - })); + let result = AzureKeyVault::new( + litellm_http::Client::plain_for_test(), + Arc::new(move |name: &str| { + (name == "AZURE_KEY_VAULT_URI") + .then(|| uri.map(str::to_owned)) + .flatten() + }), + ); if missing_environment { assert!(matches!( @@ -155,7 +158,7 @@ fn new_validates_vault_environment( #[case::local("http://localhost:8080", "https://localhost/.default")] fn derives_scope_from_vault_host(#[case] uri: &str, #[case] expected: &str) { let manager = AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), uri.parse().unwrap(), Arc::new(|_: &str| None), ) @@ -184,7 +187,7 @@ async fn missing_credentials_do_not_request_vault() { fn manager_without_credentials(server: &MockServer) -> AzureKeyVault { AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), Arc::new(|name: &str| { (name == "AZURE_CREDENTIAL").then(|| "ClientSecretCredential".to_owned()) diff --git a/litellm-rust/crates/secrets-azure/tests/live.rs b/litellm-rust/crates/secrets-azure/tests/live.rs index 18306382613..429cd3013f7 100644 --- a/litellm-rust/crates/secrets-azure/tests/live.rs +++ b/litellm-rust/crates/secrets-azure/tests/live.rs @@ -10,7 +10,7 @@ use rstest::rstest; #[ignore] async fn reads_a_real_secret() { let environment = Arc::new(ProcessEnvironment); - let manager = AzureKeyVault::new(environment).unwrap(); + let manager = AzureKeyVault::new(litellm_http::Client::plain_for_test(), environment).unwrap(); let name = std::env::var("AZURE_KEY_VAULT_LIVE_SECRET_NAME").unwrap(); let secret = manager.get_secret(&name).await.unwrap().unwrap(); assert!(matches!(&secret, Secret::String(_))); diff --git a/litellm-rust/crates/secrets-cyberark/Cargo.toml b/litellm-rust/crates/secrets-cyberark/Cargo.toml index 1c280171f4c..0a91c61ade9 100644 --- a/litellm-rust/crates/secrets-cyberark/Cargo.toml +++ b/litellm-rust/crates/secrets-cyberark/Cargo.toml @@ -8,6 +8,7 @@ repository.workspace = true [dependencies] litellm-secrets-types.workspace = true litellm-core-utils.workspace = true +litellm-http.workspace = true base64.workspace = true moka.workspace = true reqwest.workspace = true @@ -19,6 +20,7 @@ percent-encoding = "2.3" tokio = { workspace = true, features = ["sync"] } [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } rcgen = "0.14.10" rstest.workspace = true tempfile = "3.27.0" diff --git a/litellm-rust/crates/secrets-cyberark/src/error.rs b/litellm-rust/crates/secrets-cyberark/src/error.rs index 3dfeb95fe26..4f21647225b 100644 --- a/litellm-rust/crates/secrets-cyberark/src/error.rs +++ b/litellm-rust/crates/secrets-cyberark/src/error.rs @@ -18,6 +18,8 @@ pub enum Error { MissingCredentials, #[error("CyberArk client certificate could not be loaded")] ClientCertificate, + #[error("CyberArk Conjur HTTP client could not be built")] + Client(#[redact] Box), #[error("invalid refresh interval")] RefreshInterval, #[error("invalid CyberArk Conjur endpoint")] diff --git a/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs b/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs index 252a99c917f..44473c9fe90 100644 --- a/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs @@ -2,10 +2,13 @@ mod client; mod read; mod write; -use std::{fs, sync::Arc, time::Duration}; +use std::{sync::Arc, time::Duration}; use base64::{Engine, engine::general_purpose::STANDARD}; use litellm_core_utils::settings::Lookup; +use litellm_http::{ + Client, ClientIdentity, ClientVariant, HttpClientConfig, HttpClientPool, TlsSource, Verify, +}; use litellm_secrets_types::{ BaseSecretManager, CyberarkOperationContext, RotationError, SecretCache, SecretDeleter, SecretRotator, SecretValue, SecretWriteContext, SecretWriter, async_rotate_secret, @@ -37,7 +40,7 @@ const SECRET_NAME_SAFE: &AsciiSet = &NON_ALPHANUMERIC #[derive(Clone)] pub struct CyberArkSecretManager { - client: reqwest::Client, + client: Client, endpoint: reqwest::Url, account: String, username: String, @@ -45,6 +48,7 @@ pub struct CyberArkSecretManager { token: Cache<(), SecretValue>, secrets: SecretCache, authentication_lock: Arc>, + policy_load_lock: Arc>, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] diff --git a/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs b/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs index 1d99fe474d5..4b6bcfb28b0 100644 --- a/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs +++ b/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs @@ -2,7 +2,7 @@ use super::*; impl CyberArkSecretManager { pub fn with_client( - client: reqwest::Client, + client: Client, endpoint: reqwest::Url, account: String, username: String, @@ -26,10 +26,13 @@ impl CyberArkSecretManager { token, secrets, authentication_lock: Arc::new(tokio::sync::Mutex::new(())), + policy_load_lock: Arc::new(tokio::sync::Mutex::new(())), } } pub fn new( + pool: &HttpClientPool, + config: &HttpClientConfig, environment: Arc, enterprise_enabled: bool, ) -> Result { @@ -46,21 +49,34 @@ impl CyberArkSecretManager { .get(CYBERARK_SSL_VERIFY) .map(|value| !value.trim().eq_ignore_ascii_case("false")) .unwrap_or(true); - let mut builder = reqwest::Client::builder(); if !verify { litellm_tracing::warn!( "CyberArk SSL verification is disabled. This is insecure and should only be used for testing with self-signed certificates." ); - builder = builder.danger_accept_invalid_certs(true); } - if !cert.is_empty() && !key.is_empty() { - let certificate = fs::read(cert).map_err(|_| Error::ClientCertificate)?; - let private_key = fs::read(key).map_err(|_| Error::ClientCertificate)?; - let identity = reqwest::Identity::from_pem(&[certificate, private_key].concat()) - .map_err(|_| Error::ClientCertificate)?; - builder = builder.identity(identity); - } - let client = builder.build()?; + let config = HttpClientConfig { + verify: effective_verify(verify, &config.verify), + client_certificate: (!cert.is_empty() && !key.is_empty()).then(|| { + ClientIdentity::Split { + certificate: cert.into(), + key: key.into(), + } + }), + ..config.clone() + }; + let client = + pool.client(&config, ClientVariant::Provider) + .map_err(|error| match error { + litellm_http::Error::Read { + tls_source: TlsSource::ClientIdentity, + .. + } + | litellm_http::Error::InvalidPem { + tls_source: TlsSource::ClientIdentity, + .. + } => Error::ClientCertificate, + other => Error::Client(Box::new(other)), + })?; let endpoint = reqwest::Url::parse( &environment .get(CYBERARK_API_BASE) @@ -139,9 +155,35 @@ impl CyberArkSecretManager { } } +fn effective_verify(cyberark_verify: bool, host: &Verify) -> Verify { + match (cyberark_verify, host) { + (false, _) => Verify::Disabled, + (true, Verify::Disabled) => Verify::BuiltInRoots, + (true, host) => host.clone(), + } +} + fn normalize_endpoint(mut endpoint: reqwest::Url) -> reqwest::Url { if !endpoint.path().ends_with('/') { endpoint.set_path(&format!("{}/", endpoint.path())); } endpoint } + +#[cfg(test)] +mod tests { + use std::path::PathBuf; + + use super::*; + + #[test] + fn cyberark_verification_does_not_follow_a_host_that_disabled_it() { + let bundle = Verify::CaBundle(PathBuf::from("/ca.pem")); + assert_eq!( + effective_verify(true, &Verify::Disabled), + Verify::BuiltInRoots + ); + assert_eq!(effective_verify(true, &bundle), bundle); + assert_eq!(effective_verify(false, &bundle), Verify::Disabled); + } +} diff --git a/litellm-rust/crates/secrets-cyberark/src/secret_manager/write.rs b/litellm-rust/crates/secrets-cyberark/src/secret_manager/write.rs index 265c5fc6b28..87f6cea3619 100644 --- a/litellm-rust/crates/secrets-cyberark/src/secret_manager/write.rs +++ b/litellm-rust/crates/secrets-cyberark/src/secret_manager/write.rs @@ -1,5 +1,8 @@ use super::*; +const POLICY_LOAD_ATTEMPTS: u32 = 5; +const POLICY_LOAD_RETRY_DELAY: std::time::Duration = std::time::Duration::from_millis(200); + impl CyberArkSecretManager { pub async fn async_write_secret( &self, @@ -105,37 +108,43 @@ impl CyberArkSecretManager { "- !variable {}\n", serde_json::to_string(name).expect("serializing a string cannot fail") ); - let response = with_timeout( - self.client - .post(policy_url) - .header("Authorization", authorization) - .header("Content-Type", "application/x-yaml") - .body(body), - context, - ) - .send() - .await; - match response { - Ok(response) if response.status().is_success() => {} - Ok(response) - if matches!( - response.status(), - reqwest::StatusCode::CONFLICT | reqwest::StatusCode::UNPROCESSABLE_ENTITY - ) => - { - litellm_tracing::debug!( - "CyberArk variable policy already exists or conflicts: {}", - response.status() - ); - } - Ok(response) => { - litellm_tracing::warn!( - "Could not ensure CyberArk variable exists: {}", - response.status() - ); - } - Err(error) => { - litellm_tracing::warn!("Error ensuring CyberArk variable exists: {error}"); + let _policy_load = self.policy_load_lock.lock().await; + for attempt in 0..POLICY_LOAD_ATTEMPTS { + let response = with_timeout( + self.client + .post(policy_url.clone()) + .header("Authorization", authorization.clone()) + .header("Content-Type", "application/x-yaml") + .body(body.clone()), + context, + ) + .send() + .await; + match response { + Ok(response) + if response.status() == reqwest::StatusCode::CONFLICT + && attempt + 1 < POLICY_LOAD_ATTEMPTS => + { + tokio::time::sleep(POLICY_LOAD_RETRY_DELAY * 2_u32.pow(attempt)).await; + } + Ok(response) if response.status().is_success() => return, + Ok(response) if response.status() == reqwest::StatusCode::UNPROCESSABLE_ENTITY => { + litellm_tracing::debug!( + "CyberArk variable policy was rejected as unprocessable" + ); + return; + } + Ok(response) => { + litellm_tracing::warn!( + "Could not ensure CyberArk variable exists: {}", + response.status() + ); + return; + } + Err(error) => { + litellm_tracing::warn!("Error ensuring CyberArk variable exists: {error}"); + return; + } } } } diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs index 2048e067b6e..783ba6ff67a 100644 --- a/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs @@ -7,6 +7,8 @@ use std::{ }; use base64::{Engine, engine::general_purpose::STANDARD}; +use litellm_core_utils::settings::Lookup; +use litellm_http::{HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver}; use litellm_secrets_cyberark::{CyberArkSecretManager, DeleteOutcome, Error}; use litellm_secrets_types::{BaseSecretManager, CyberarkOperationContext, SecretValue}; use rstest::{fixture, rstest}; diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs index fbd4317f446..1bea77e49ae 100644 --- a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs @@ -77,7 +77,7 @@ async fn authentication_encodes_login(#[case] username: &str, #[case] expected_p .mount(&server) .await; let manager = CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "acct".into(), username.into(), @@ -211,25 +211,25 @@ fn new_validates_credentials_before_license_and_configuration() { let empty: Arc = Arc::new(|_: &str| None); assert!(matches!( - CyberArkSecretManager::new(empty, true), + from_environment(empty, true), Err(Error::MissingCredentials) )); assert!(matches!( - CyberArkSecretManager::new( + from_environment( Arc::new(|name: &str| (name == "CYBERARK_API_KEY").then(|| "k3y".into())), false ), Err(Error::EnterpriseRequired) )); assert!(matches!( - CyberArkSecretManager::new( + from_environment( Arc::new(|name: &str| (name == "CYBERARK_CLIENT_CERT").then(|| "cert".into())), true ), Err(Error::MissingCredentials) )); assert!(matches!( - CyberArkSecretManager::new( + from_environment( Arc::new(|name: &str| match name { "CYBERARK_API_KEY" => Some("k3y".into()), "CYBERARK_REFRESH_INTERVAL" => Some("abc".into()), @@ -240,7 +240,7 @@ fn new_validates_credentials_before_license_and_configuration() { Err(Error::RefreshInterval) )); assert!(matches!( - CyberArkSecretManager::new( + from_environment( Arc::new(|name: &str| match name { "CYBERARK_API_KEY" => Some("k3y".into()), "CYBERARK_API_BASE" => Some("not a url".into()), @@ -254,7 +254,7 @@ fn new_validates_credentials_before_license_and_configuration() { #[rstest] fn certificate_only_credentials_are_validated_as_a_client_identity() { - let result = CyberArkSecretManager::new( + let result = from_environment( Arc::new(|name: &str| match name { "CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()), "CYBERARK_CLIENT_KEY" => Some("/missing/key".into()), @@ -295,7 +295,7 @@ async fn configured_client_identity_preserves_auth_request_and_read_result( let endpoint = server.uri(); let certificate = client_identity_directory.path().join("client.crt"); let key = client_identity_directory.path().join("client.key"); - let manager = CyberArkSecretManager::new( + let manager = from_environment( Arc::new(move |name: &str| match name { "CYBERARK_API_BASE" => Some(endpoint.clone()), "CYBERARK_API_KEY" => Some(api_key.into()), @@ -337,7 +337,7 @@ fn invalid_client_identity_is_not_ignored( let certificate = client_identity_directory.path().join("client.crt"); let key = client_identity_directory.path().join("client.key"); - let result = CyberArkSecretManager::new( + let result = from_environment( Arc::new(move |name: &str| match name { "CYBERARK_API_KEY" => Some(api_key.into()), "CYBERARK_CLIENT_CERT" => Some(certificate.to_str().unwrap().into()), @@ -354,7 +354,7 @@ fn invalid_client_identity_is_not_ignored( #[case::certificate_only("")] #[case::certificate_and_api_key("k3y")] fn client_identity_does_not_bypass_the_enterprise_requirement(#[case] api_key: &'static str) { - let result = CyberArkSecretManager::new( + let result = from_environment( Arc::new(move |name: &str| match name { "CYBERARK_API_KEY" => Some(api_key.into()), "CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()), @@ -381,7 +381,7 @@ async fn new_reads_environment_defaults_end_to_end() { .mount(&server) .await; let endpoint = server.uri(); - let manager = CyberArkSecretManager::new( + let manager = from_environment( Arc::new(move |name: &str| match name { "CYBERARK_API_BASE" => Some(endpoint.clone()), "CYBERARK_API_KEY" => Some("k3y".into()), @@ -404,7 +404,7 @@ async fn new_reads_environment_defaults_end_to_end() { #[rstest] fn new_reports_missing_client_certificate_files() { assert!(matches!( - CyberArkSecretManager::new( + from_environment( Arc::new(|name: &str| match name { "CYBERARK_API_KEY" => Some("k3y".into()), "CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()), @@ -433,7 +433,7 @@ async fn trailing_slash_endpoint_preserves_base_path() { .await; let endpoint = format!("{}/prefix/", server.uri()).parse().unwrap(); let manager = CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), endpoint, "acct".into(), "admin".into(), diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs index f5ba7a63273..6fa15ecba9b 100644 --- a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs @@ -48,9 +48,21 @@ pub(super) fn client_identity_directory() -> tempfile::TempDir { directory } +pub(super) fn from_environment( + environment: Arc, + enterprise_enabled: bool, +) -> Result { + CyberArkSecretManager::new( + &HttpClientPool::new(Arc::new(PublicDnsResolver)), + &Resolution::from(&HttpSettings::default()).config, + environment, + enterprise_enabled, + ) +} + pub(super) fn manager(server: &MockServer, ttl: Duration) -> CyberArkSecretManager { CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "acct".into(), "admin".into(), diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs index 331a26c6119..764a591bc71 100644 --- a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs @@ -6,7 +6,7 @@ async fn rejected_write_token_is_reauthenticated_once() { let server = MockServer::start().await; mount_auth(&server, 2).await; Mock::given(path("/policies/acct/policy/root")) - .respond_with(ResponseTemplate::new(409)) + .respond_with(ResponseTemplate::new(201)) .expect(1) .mount(&server) .await; @@ -45,7 +45,6 @@ async fn rejected_write_token_is_reauthenticated_once() { #[rstest] #[case::created(201)] -#[case::already_exists(409)] #[case::unprocessable(422)] #[case::server_error(500)] #[tokio::test] @@ -81,13 +80,81 @@ async fn writes_tolerate_policy_status_and_cache_value(#[case] policy_status: u1 ); } +#[rstest] +#[tokio::test] +async fn policy_load_conflict_is_retried_before_the_value_write() { + let server = MockServer::start().await; + mount_auth(&server, 1).await; + let policy_loads = Arc::new(AtomicUsize::new(0)); + let policy_loads_for_response = Arc::clone(&policy_loads); + Mock::given(path("/policies/acct/policy/root")) + .respond_with(move |_: &Request| { + if policy_loads_for_response.fetch_add(1, Ordering::SeqCst) < 2 { + ResponseTemplate::new(409) + } else { + ResponseTemplate::new(201) + } + }) + .expect(3) + .mount(&server) + .await; + let policy_loads_at_value_write = Arc::clone(&policy_loads); + Mock::given(method("POST")) + .and(path("/secrets/acct/variable/key")) + .respond_with(move |_: &Request| { + if policy_loads_at_value_write.load(Ordering::SeqCst) == 3 { + ResponseTemplate::new(201) + } else { + ResponseTemplate::new(404) + } + }) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + + manager + .async_write_secret("key", &SecretValue::new("v"), None) + .await + .unwrap(); +} + +#[rstest] +#[tokio::test] +async fn concurrent_writes_load_policy_one_at_a_time() { + let server = MockServer::start().await; + mount_auth(&server, 1).await; + Mock::given(path("/policies/acct/policy/root")) + .respond_with(ResponseTemplate::new(201).set_delay(Duration::from_millis(100))) + .expect(4) + .mount(&server) + .await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(201)) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + let started = std::time::Instant::now(); + + let value = SecretValue::new("v"); + let results = tokio::join!( + manager.async_write_secret("key-0", &value, None), + manager.async_write_secret("key-1", &value, None), + manager.async_write_secret("key-2", &value, None), + manager.async_write_secret("key-3", &value, None), + ); + + assert!(results.0.is_ok() && results.1.is_ok() && results.2.is_ok() && results.3.is_ok()); + assert!(started.elapsed() >= Duration::from_millis(400)); +} + #[rstest] #[tokio::test] async fn failed_value_write_is_not_cached() { let server = MockServer::start().await; mount_auth(&server, 1).await; Mock::given(path("/policies/acct/policy/root")) - .respond_with(ResponseTemplate::new(409)) + .respond_with(ResponseTemplate::new(201)) .mount(&server) .await; Mock::given(path("/secrets/acct/variable/key")) @@ -166,7 +233,7 @@ async fn writes_match_python_parity_fixture(parity_fixture: ParityFixture) { .mount(&server) .await; let manager = CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), parity_fixture.account, parity_fixture.username, @@ -230,7 +297,7 @@ async fn live_conjur_round_trip() { .as_nanos() ); let manager = CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), endpoint.clone(), account.clone(), username.clone(), @@ -245,7 +312,7 @@ async fn live_conjur_round_trip() { .await .unwrap(); let verifier = CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), endpoint.clone(), account.clone(), username.clone(), diff --git a/litellm-rust/crates/secrets-google/Cargo.toml b/litellm-rust/crates/secrets-google/Cargo.toml index 805eb80740d..208b5ddd03f 100644 --- a/litellm-rust/crates/secrets-google/Cargo.toml +++ b/litellm-rust/crates/secrets-google/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true moka.workspace = true tokio.workspace = true litellm-auth-gcp = { workspace = true, features = ["google-sdk"] } @@ -24,6 +25,7 @@ serde.workspace = true reqwest.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } google-cloud-auth.workspace = true rstest.workspace = true wiremock = "0.6.5" diff --git a/litellm-rust/crates/secrets-google/src/secret_manager.rs b/litellm-rust/crates/secrets-google/src/secret_manager.rs index b8787999e12..08ba466b799 100644 --- a/litellm-rust/crates/secrets-google/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-google/src/secret_manager.rs @@ -23,7 +23,7 @@ const CACHE_CAPACITY: u64 = 200; #[derive(Clone)] pub struct GoogleSecretManager { - client: reqwest::Client, + client: litellm_http::Client, credentials: Arc, endpoint: reqwest::Url, project: String, @@ -46,7 +46,7 @@ struct Payload { impl GoogleSecretManager { pub fn with_client( - client: reqwest::Client, + client: litellm_http::Client, endpoint: reqwest::Url, project: String, environment: Arc, @@ -79,6 +79,7 @@ impl GoogleSecretManager { } pub fn new( + client: litellm_http::Client, environment: Arc, enterprise_enabled: bool, ) -> Result { @@ -104,7 +105,7 @@ impl GoogleSecretManager { .get(GOOGLE_SECRET_MANAGER_ALWAYS_READ_SECRET_MANAGER) .is_some_and(|v| v.eq_ignore_ascii_case("true")); Self::with_client( - reqwest::Client::new(), + client, reqwest::Url::parse("https://secretmanager.googleapis.com").expect("static URL"), project, environment, diff --git a/litellm-rust/crates/secrets-google/tests/secret_manager.rs b/litellm-rust/crates/secrets-google/tests/secret_manager.rs index 0d7efc4b1b3..e9bee7633f5 100644 --- a/litellm-rust/crates/secrets-google/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-google/tests/secret_manager.rs @@ -11,7 +11,7 @@ use wiremock::{ fn manager(server: &MockServer, always_read: bool, ttl: Duration) -> GoogleSecretManager { GoogleSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "project".into(), Arc::new(|name: &str| (name == "VERTEX_AI_API_KEY").then(|| "token".into())), @@ -214,11 +214,19 @@ async fn always_read_and_expired_cache_fetch_again( #[rstest] fn google_manager_requires_host_license_and_project_configuration() { assert!(matches!( - GoogleSecretManager::new(Arc::new(|_: &str| None), false), + GoogleSecretManager::new( + litellm_http::Client::plain_for_test(), + Arc::new(|_: &str| None), + false + ), Err(Error::EnterpriseRequired) )); assert!(matches!( - GoogleSecretManager::new(Arc::new(|_: &str| None), true), + GoogleSecretManager::new( + litellm_http::Client::plain_for_test(), + Arc::new(|_: &str| None), + true + ), Err(Error::MissingEnvironment( "GOOGLE_SECRET_MANAGER_PROJECT_ID" )) @@ -236,7 +244,7 @@ fn google_manager_rejects_invalid_refresh_intervals(#[case] variable: &'static s }); assert!(matches!( - GoogleSecretManager::new(environment, true), + GoogleSecretManager::new(litellm_http::Client::plain_for_test(), environment, true), Err(Error::RefreshInterval) )); } diff --git a/litellm-rust/crates/secrets-types/Cargo.toml b/litellm-rust/crates/secrets-types/Cargo.toml index dcd06d1a741..6dd847ec989 100644 --- a/litellm-rust/crates/secrets-types/Cargo.toml +++ b/litellm-rust/crates/secrets-types/Cargo.toml @@ -11,6 +11,7 @@ moka.workspace = true tokio = { workspace = true, features = ["sync"] } serde.workspace = true serde_json.workspace = true +strum.workspace = true thiserror.workspace = true veil.workspace = true diff --git a/litellm-rust/crates/secrets-types/src/config.rs b/litellm-rust/crates/secrets-types/src/config.rs index 44acf512224..82d48f7b2e1 100644 --- a/litellm-rust/crates/secrets-types/src/config.rs +++ b/litellm-rust/crates/secrets-types/src/config.rs @@ -1,11 +1,13 @@ use std::collections::BTreeMap; use serde::{Deserialize, Serialize}; +use strum::IntoStaticStr; use crate::SecretValue; -#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)] +#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, IntoStaticStr, PartialEq, Serialize)] #[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] pub enum KeyManagementSystem { GoogleKms, AzureKeyVault, diff --git a/litellm-rust/crates/secrets/Cargo.toml b/litellm-rust/crates/secrets/Cargo.toml index 17acc01682b..f855a8a64a6 100644 --- a/litellm-rust/crates/secrets/Cargo.toml +++ b/litellm-rust/crates/secrets/Cargo.toml @@ -15,7 +15,7 @@ cyberark = ["dep:litellm-secrets-cyberark"] [dependencies] futures-util.workspace = true -litellm-python-compat = { path = "../python-compat" } +litellm-python-compat.workspace = true litellm-secrets-types.workspace = true litellm-secrets-aws = { workspace = true, optional = true } litellm-secrets-google = { workspace = true, optional = true } @@ -23,6 +23,7 @@ litellm-secrets-hashicorp = { workspace = true, optional = true } litellm-secrets-azure = { workspace = true, optional = true } litellm-secrets-cyberark = { workspace = true, optional = true } litellm-core-utils.workspace = true +litellm-http.workspace = true base64.workspace = true serde.workspace = true strum.workspace = true @@ -33,6 +34,7 @@ moka.workspace = true tokio = { workspace = true, features = ["fs"] } [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } rstest.workspace = true wiremock = "0.6.5" tempfile = "3" diff --git a/litellm-rust/crates/secrets/src/error.rs b/litellm-rust/crates/secrets/src/error.rs index 07f2f205bec..7fc756c5c7d 100644 --- a/litellm-rust/crates/secrets/src/error.rs +++ b/litellm-rust/crates/secrets/src/error.rs @@ -10,6 +10,8 @@ pub enum Error { InvalidCiphertext, #[error("decrypted value is not UTF-8")] Utf8, + #[error(transparent)] + Client(#[from] litellm_http::Error), #[error("unsupported OIDC provider or missing build feature")] UnsupportedOidc, #[error("OIDC reference requires a provider and audience")] diff --git a/litellm-rust/crates/secrets/src/native.rs b/litellm-rust/crates/secrets/src/native.rs index 80f0e46245c..8edcebe7b12 100644 --- a/litellm-rust/crates/secrets/src/native.rs +++ b/litellm-rust/crates/secrets/src/native.rs @@ -1,10 +1,13 @@ use std::sync::Arc; use litellm_core_utils::settings::Lookup; +use litellm_http::{HttpClientConfig, HttpClientPool}; use crate::{Error, KeyManagementSettings, KeyManagementSystem, SecretManager}; pub async fn load_native_manager( + _pool: &HttpClientPool, + _config: &HttpClientConfig, system: KeyManagementSystem, settings: KeyManagementSettings, environment: Arc, @@ -29,14 +32,19 @@ pub async fn load_native_manager( } #[cfg(feature = "azure")] (KeyManagementSystem::AzureKeyVault, _, environment, _) => Ok( - SecretManager::AzureKeyVault(crate::azure::AzureKeyVault::new(environment)?), + SecretManager::AzureKeyVault(crate::azure::AzureKeyVault::new( + _pool.client(_config, litellm_http::ClientVariant::Provider)?, + environment, + )?), ), #[cfg(feature = "google")] - (KeyManagementSystem::GoogleSecretManager, _, environment, enterprise_enabled) => { - Ok(SecretManager::GoogleSecretManager( - crate::google::GoogleSecretManager::new(environment, enterprise_enabled)?, - )) - } + (KeyManagementSystem::GoogleSecretManager, _, environment, enterprise_enabled) => Ok( + SecretManager::GoogleSecretManager(crate::google::GoogleSecretManager::new( + _pool.client(_config, litellm_http::ClientVariant::Provider)?, + environment, + enterprise_enabled, + )?), + ), #[cfg(feature = "google")] (KeyManagementSystem::GoogleKms, _, environment, _) => { crate::google::load_google_kms(Some(true), environment) @@ -51,11 +59,14 @@ pub async fn load_native_manager( )) } #[cfg(feature = "cyberark")] - (KeyManagementSystem::Cyberark, _, environment, enterprise_enabled) => { - Ok(SecretManager::Cyberark( - crate::cyberark::CyberArkSecretManager::new(environment, enterprise_enabled)?, - )) - } + (KeyManagementSystem::Cyberark, _, environment, enterprise_enabled) => Ok( + SecretManager::Cyberark(crate::cyberark::CyberArkSecretManager::new( + _pool, + _config, + environment, + enterprise_enabled, + )?), + ), _ => Err(Error::NativeBackendUnavailable), } } diff --git a/litellm-rust/crates/secrets/src/oidc.rs b/litellm-rust/crates/secrets/src/oidc.rs index f3c1e38ce7b..b6e8dbc123b 100644 --- a/litellm-rust/crates/secrets/src/oidc.rs +++ b/litellm-rust/crates/secrets/src/oidc.rs @@ -5,6 +5,7 @@ use std::{ use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD}; use litellm_core_utils::settings::Lookup; +use litellm_http::Client; use moka::future::Cache; use serde::Deserialize; @@ -82,7 +83,7 @@ impl NumericDate { } pub struct OidcResolver { - client: reqwest::Client, + client: Client, google_identity_endpoint: reqwest::Url, cache: Cache, clock: fn() -> SystemTime, @@ -90,25 +91,17 @@ pub struct OidcResolver { azure_token_provider: std::sync::Arc, } -impl Default for OidcResolver { - fn default() -> Self { - let client = reqwest::Client::builder() - .timeout(Duration::from_secs(600)) - .connect_timeout(Duration::from_secs(5)) - .build() - .expect("HTTP client configuration"); - Self::new( - client, - reqwest::Url::parse("http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity").expect("static URL"), - ) - } -} +const GOOGLE_IDENTITY_ENDPOINT: &str = + "http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity"; + +const REQUEST_TIMEOUT: Duration = Duration::from_secs(600); impl OidcResolver { - pub fn new(client: reqwest::Client, google_identity_endpoint: reqwest::Url) -> Self { + pub fn new(client: Client) -> Self { Self { client, - google_identity_endpoint, + google_identity_endpoint: reqwest::Url::parse(GOOGLE_IDENTITY_ENDPOINT) + .expect("static URL"), cache: Cache::builder() .max_capacity(200) .time_to_live(GOOGLE_TOKEN_MAX_TTL) @@ -121,6 +114,13 @@ impl OidcResolver { } } + pub fn with_google_identity_endpoint(self, google_identity_endpoint: reqwest::Url) -> Self { + Self { + google_identity_endpoint, + ..self + } + } + #[cfg(feature = "azure")] pub fn with_azure_token_provider( self, @@ -180,6 +180,7 @@ impl OidcResolver { let response = self .client .get(url) + .timeout(REQUEST_TIMEOUT) .query(&[("audience", audience)]) .bearer_auth(authorization) .header("Accept", "application/json; api-version=2.0") @@ -214,6 +215,7 @@ impl OidcResolver { let response = self .client .get(self.google_identity_endpoint.clone()) + .timeout(REQUEST_TIMEOUT) .query(&[("audience", audience)]) .header("Metadata-Flavor", "Google") .send() diff --git a/litellm-rust/crates/secrets/src/resolver.rs b/litellm-rust/crates/secrets/src/resolver.rs index 69445e5b410..830e7c22cd5 100644 --- a/litellm-rust/crates/secrets/src/resolver.rs +++ b/litellm-rust/crates/secrets/src/resolver.rs @@ -1,10 +1,7 @@ use std::sync::Arc; use crate::compatibility::python_manager_string; -use litellm_core_utils::{ - serde_compat::parse_str_bool, - settings::{Lookup, ProcessEnvironment}, -}; +use litellm_core_utils::{serde_compat::parse_str_bool, settings::Lookup}; use crate::state::{LookupTarget, normalize_secret_name}; use crate::{Error, OidcResolver, Secret, SecretManagerState, SecretValue}; @@ -24,16 +21,6 @@ pub struct SecretResolver { python_compatible: bool, } -impl Default for SecretResolver { - fn default() -> Self { - Self::new( - Arc::new(SecretManagerState::default()), - Arc::new(ProcessEnvironment), - OidcResolver::default(), - ) - } -} - impl SecretResolver { pub fn new( state: Arc, diff --git a/litellm-rust/crates/secrets/src/source.rs b/litellm-rust/crates/secrets/src/source.rs index a1615bad055..3c86fd25e9a 100644 --- a/litellm-rust/crates/secrets/src/source.rs +++ b/litellm-rust/crates/secrets/src/source.rs @@ -37,15 +37,14 @@ impl SecretSource for SecretResolver { } } -#[derive(Default)] pub struct EnvironmentSecrets(SecretResolver); impl EnvironmentSecrets { - pub fn python_compatible() -> Self { + pub fn python_compatible(client: litellm_http::Client) -> Self { Self(SecretResolver::new_python_compatible( Arc::new(crate::SecretManagerState::default()), Arc::new(litellm_core_utils::settings::ProcessEnvironment), - crate::OidcResolver::default(), + crate::OidcResolver::new(client), )) } } diff --git a/litellm-rust/crates/secrets/tests/aws.rs b/litellm-rust/crates/secrets/tests/aws.rs index 174d0881339..910b1855b77 100644 --- a/litellm-rust/crates/secrets/tests/aws.rs +++ b/litellm-rust/crates/secrets/tests/aws.rs @@ -53,7 +53,7 @@ async fn read_results_follow_the_selected_failure_policy( }, )), Arc::new(move |_: &str| environment.map(str::to_owned)), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ) .with_failure_policy(policy); let result = resolver @@ -97,7 +97,7 @@ async fn primary_secret_values_other_than_strings_resolve_to_none( let resolver = SecretResolver::new_python_compatible( Arc::new(state(&server, settings)), Arc::new(|_: &str| Some("fallback".into())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); let text = value.as_str(); assert_eq!( @@ -157,7 +157,7 @@ async fn gating_prediction_matches_actual_lookup( let resolver = SecretResolver::new_python_compatible( Arc::new(state), Arc::new(|_: &str| Some("environment".into())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver diff --git a/litellm-rust/crates/secrets/tests/azure.rs b/litellm-rust/crates/secrets/tests/azure.rs index b844b198cd4..60f165add54 100644 --- a/litellm-rust/crates/secrets/tests/azure.rs +++ b/litellm-rust/crates/secrets/tests/azure.rs @@ -22,7 +22,7 @@ async fn azure_handler_reads_missing_and_failed_secrets() { .await; let manager = SecretManager::AzureKeyVault( AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), std::sync::Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "fake".to_owned())), ) @@ -81,7 +81,7 @@ async fn successful_azure_responses_do_not_fall_back_when_the_value_is_empty_or_ .mount(&server) .await; let manager = AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "token".into())), ) @@ -92,7 +92,7 @@ async fn successful_azure_responses_do_not_fall_back_when_the_value_is_empty_or_ Default::default(), )), Arc::new(|_: &str| Some("environment".into())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver diff --git a/litellm-rust/crates/secrets/tests/common_read_contract.rs b/litellm-rust/crates/secrets/tests/common_read_contract.rs index 698e1ad8f63..2e8fbf46008 100644 --- a/litellm-rust/crates/secrets/tests/common_read_contract.rs +++ b/litellm-rust/crates/secrets/tests/common_read_contract.rs @@ -59,7 +59,7 @@ fn manager(provider: Provider, server: &MockServer) -> SecretManager { ), Provider::Azure => SecretManager::AzureKeyVault( AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), environment, ) @@ -67,7 +67,7 @@ fn manager(provider: Provider, server: &MockServer) -> SecretManager { ), Provider::Google => SecretManager::GoogleSecretManager( GoogleSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "project".into(), environment, @@ -84,7 +84,7 @@ fn manager(provider: Provider, server: &MockServer) -> SecretManager { .unwrap(), ), Provider::Cyberark => SecretManager::Cyberark(CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "acct".into(), "admin".into(), @@ -240,7 +240,7 @@ async fn python_read_failures_preserve_provider_fallback_rules( KeyManagementSettings::default(), )), Arc::new(move |_: &str| environment_value.map(str::to_owned)), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); let expected = if matches!(provider, Provider::Aws) { None diff --git a/litellm-rust/crates/secrets/tests/cyberark.rs b/litellm-rust/crates/secrets/tests/cyberark.rs index 706c35752d7..b94bd9dc0ad 100644 --- a/litellm-rust/crates/secrets/tests/cyberark.rs +++ b/litellm-rust/crates/secrets/tests/cyberark.rs @@ -24,7 +24,7 @@ async fn cyberark_handler_reads_values_and_surfaces_errors() { .mount(&server) .await; let manager = SecretManager::Cyberark(CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "acct".into(), "admin".into(), diff --git a/litellm-rust/crates/secrets/tests/google.rs b/litellm-rust/crates/secrets/tests/google.rs index 67954fedb78..0b45a90c10b 100644 --- a/litellm-rust/crates/secrets/tests/google.rs +++ b/litellm-rust/crates/secrets/tests/google.rs @@ -25,7 +25,7 @@ async fn google_resolver_distinguishes_absence_from_failure(#[case] status: u16) _ => None, }); let manager = GoogleSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "project".into(), environment.clone(), @@ -40,7 +40,7 @@ async fn google_resolver_distinguishes_absence_from_failure(#[case] status: u16) let resolver = SecretResolver::new_python_compatible( Arc::new(state), environment, - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ) .with_failure_policy(FailurePolicy::Propagate); let result = resolver.get_secret_str("KEY", None).await; diff --git a/litellm-rust/crates/secrets/tests/hashicorp.rs b/litellm-rust/crates/secrets/tests/hashicorp.rs index bc35b88018e..e10e903c25a 100644 --- a/litellm-rust/crates/secrets/tests/hashicorp.rs +++ b/litellm-rust/crates/secrets/tests/hashicorp.rs @@ -56,7 +56,7 @@ async fn hashicorp_handler_resolves_found_missing_and_failed_values() { }, )), Arc::new(|_: &str| None), - litellm_secrets::OidcResolver::default(), + litellm_secrets::OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( found_resolver @@ -131,7 +131,7 @@ async fn hashicorp_handler_resolves_found_missing_and_failed_values() { let failed_resolver = SecretResolver::new_python_compatible( Arc::new(failed_state), Arc::new(|_: &str| None), - litellm_secrets::OidcResolver::default(), + litellm_secrets::OidcResolver::new(litellm_http::Client::plain_for_test()), ) .with_failure_policy(FailurePolicy::Propagate); assert!(matches!( diff --git a/litellm-rust/crates/secrets/tests/oidc.rs b/litellm-rust/crates/secrets/tests/oidc.rs index afc49e8231d..6f72bd9d645 100644 --- a/litellm-rust/crates/secrets/tests/oidc.rs +++ b/litellm-rust/crates/secrets/tests/oidc.rs @@ -30,7 +30,7 @@ async fn environment_sources_resolve_expected_value( ("CIRCLE_OIDC_TOKEN_V2", "circle-v2"), ]); assert_eq!( - OidcResolver::default() + OidcResolver::new(litellm_http::Client::plain_for_test()) .resolve(reference, env.as_ref()) .await .unwrap() @@ -43,7 +43,7 @@ async fn environment_sources_resolve_expected_value( #[tokio::test] async fn environment_sources_bypass_boolean_conversion_and_defaults() { let env = environment(&[("TOKEN", "true")]); - let oidc = OidcResolver::default(); + let oidc = OidcResolver::new(litellm_http::Client::plain_for_test()); let resolver = SecretResolver::new(Arc::new(SecretManagerState::default()), env, oidc); assert_eq!( resolver @@ -94,7 +94,7 @@ async fn github_requests_are_authenticated_cached_and_revalidate_environment() { ), ("ACTIONS_ID_TOKEN_REQUEST_TOKEN", "request-token"), ]); - let oidc = OidcResolver::default(); + let oidc = OidcResolver::new(litellm_http::Client::plain_for_test()); for _ in 0..2 { assert_eq!( oidc.resolve("oidc/github/https://service/oidc/path", env.as_ref()) @@ -131,7 +131,7 @@ async fn file_allowlist_resolves_symlinks_while_environment_paths_remain_explici ("PATH_TOKEN", private.to_str().unwrap()), ("AZURE_FEDERATED_TOKEN_FILE", token.to_str().unwrap()), ]); - let oidc = OidcResolver::default(); + let oidc = OidcResolver::new(litellm_http::Client::plain_for_test()); assert_eq!( oidc.resolve(&format!("oidc/file/{}", token.display()), env.as_ref()) .await @@ -213,8 +213,9 @@ async fn google_expiry_caps_cache_and_preserves_audience( .expect(calls) .mount(&server) .await; - let oidc = - OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()).with_clock(now); + let oidc = OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_google_identity_endpoint(server.uri().parse().unwrap()) + .with_clock(now); for _ in 0..2 { assert_eq!( oidc.resolve( @@ -234,7 +235,7 @@ async fn google_expiry_caps_cache_and_preserves_audience( #[tokio::test] async fn google_oidc_requires_its_build_feature() { assert!(matches!( - OidcResolver::default() + OidcResolver::new(litellm_http::Client::plain_for_test()) .resolve("oidc/google/audience", environment(&[]).as_ref()) .await, Err(Error::UnsupportedOidc) @@ -245,7 +246,7 @@ async fn google_oidc_requires_its_build_feature() { #[tokio::test] async fn azure_oidc_without_a_token_file_requires_its_build_feature() { assert!(matches!( - OidcResolver::default() + OidcResolver::new(litellm_http::Client::plain_for_test()) .resolve("oidc/azure/scope", environment(&[]).as_ref()) .await, Err(Error::UnsupportedOidc) @@ -261,7 +262,7 @@ async fn invalid_references_fail_before_environment_lookup( #[case] reference: &str, #[case] unsupported: bool, ) { - let error = OidcResolver::default() + let error = OidcResolver::new(litellm_http::Client::plain_for_test()) .resolve(reference, &|_: &str| { panic!("invalid reference reached environment lookup") }) @@ -283,7 +284,8 @@ async fn unreadable_expiry_keeps_python_cache_fallback(#[case] token: &str) { .expect(1) .mount(&server) .await; - let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()); + let resolver = OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_google_identity_endpoint(server.uri().parse().unwrap()); for _ in 0..2 { assert_eq!( resolver @@ -334,7 +336,8 @@ async fn azure_oidc_acquires_the_requested_scope_and_preserves_failures(#[case] }) } } - let oidc = OidcResolver::default().with_azure_token_provider(Arc::new(Provider(failed))); + let oidc = OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_azure_token_provider(Arc::new(Provider(failed))); let resolver = SecretResolver::new_python_compatible( Arc::new(SecretManagerState::default()), environment(&[("AZURE_CLIENT_ID", "client-id")]), @@ -361,7 +364,7 @@ async fn azure_oidc_acquires_the_requested_scope_and_preserves_failures(#[case] #[tokio::test] async fn missing_oidc_environment_is_an_error(#[case] reference: &str) { assert!(matches!( - OidcResolver::default() + OidcResolver::new(litellm_http::Client::plain_for_test()) .resolve(reference, environment(&[]).as_ref()) .await, Err(Error::MissingEnvironment) @@ -380,7 +383,8 @@ async fn google_oidc_failures_are_not_cached_or_hidden_by_defaults() { let resolver = SecretResolver::new_python_compatible( Arc::new(SecretManagerState::default()), environment(&[]), - OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()), + OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_google_identity_endpoint(server.uri().parse().unwrap()), ); for _ in 0..2 { assert!(matches!( @@ -420,7 +424,8 @@ async fn google_tokens_expire_at_the_python_cache_deadline( .expect(2) .mount(&server) .await; - let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()) + let resolver = OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_google_identity_endpoint(server.uri().parse().unwrap()) .with_clock(|| UNIX_EPOCH + Duration::from_secs(1000)); assert_eq!( resolver @@ -472,7 +477,8 @@ async fn google_cache_uses_payload_expiry_without_requiring_a_jwt_header() { .expect(2) .mount(&server) .await; - let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()) + let resolver = OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_google_identity_endpoint(server.uri().parse().unwrap()) .with_clock(|| UNIX_EPOCH + Duration::from_secs(1000)); for _ in 0..2 { assert_eq!( diff --git a/litellm-rust/crates/secrets/tests/resolution.rs b/litellm-rust/crates/secrets/tests/resolution.rs index bed762adc59..37557c36fb5 100644 --- a/litellm-rust/crates/secrets/tests/resolution.rs +++ b/litellm-rust/crates/secrets/tests/resolution.rs @@ -12,7 +12,7 @@ fn resolver(value: Option<&str>) -> SecretResolver { SecretResolver::new_python_compatible( Arc::new(SecretManagerState::default()), Arc::new(move |_: &str| value.clone()), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ) } @@ -35,7 +35,7 @@ async fn native_reads_preserve_strings_and_report_conversion_errors(#[case] mana let resolver = SecretResolver::new( Arc::new(state), Arc::new(move |_: &str| Some(raw.to_owned())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver @@ -75,7 +75,7 @@ async fn native_defaults_apply_to_absence_but_never_hide_provider_failures() { KeyManagementSettings::default(), )), Arc::new(|_: &str| None), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); let result = resolver .get_secret_str("key", Some(SecretValue::new("default"))) @@ -189,7 +189,7 @@ fn managed(reply: Result, ()>, environment: Option<&'static str>) KeyManagementSettings::default(), )), Arc::new(move |_: &str| environment.map(str::to_owned)), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ) .with_failure_policy(FailurePolicy::EnvironmentFallback) } @@ -281,7 +281,7 @@ async fn prefix_is_removed_once_and_resolved_from_environment() { let resolver = SecretResolver::new_python_compatible( Arc::new(state), Arc::new(|name: &str| (name == "os.environ/KEY").then(|| "value".into())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver @@ -340,7 +340,7 @@ async fn excluded_hosted_keys_keep_the_python_manager_conversion_path( }, )), Arc::new(move |_: &str| Some(raw.to_owned())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver.get_secret("KEY", None).await.unwrap(), @@ -373,7 +373,7 @@ async fn azure_callback_absence_preserves_none_but_errors_fall_back( KeyManagementSettings::default(), )), Arc::new(|_: &str| Some("environment".into())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver diff --git a/litellm-rust/crates/secrets/tests/source.rs b/litellm-rust/crates/secrets/tests/source.rs index b4782c6af86..17210940a67 100644 --- a/litellm-rust/crates/secrets/tests/source.rs +++ b/litellm-rust/crates/secrets/tests/source.rs @@ -15,7 +15,7 @@ mod tests { #[case] expected: Option<&str>, ) { unsafe { std::env::set_var(name, value) }; - let secret = EnvironmentSecrets::python_compatible() + let secret = EnvironmentSecrets::python_compatible(litellm_http::Client::plain_for_test()) .resolve(&[name]) .await .unwrap() @@ -42,7 +42,7 @@ async fn dynamic_names_use_the_same_resolver_and_snapshots_never_do_fresh_lookup reads.fetch_add(1, Ordering::SeqCst); (name != "missing").then(|| name.to_owned()) }), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); let snapshot = source.resolve(&["declared", "missing"]).await.unwrap(); let name = format!("runtime-{}", "key"); diff --git a/litellm-rust/crates/testkit/src/lib.rs b/litellm-rust/crates/testkit/src/lib.rs index 9ea6123a176..423adeec426 100644 --- a/litellm-rust/crates/testkit/src/lib.rs +++ b/litellm-rust/crates/testkit/src/lib.rs @@ -1,3 +1,9 @@ +#![allow( + clippy::disallowed_types, + clippy::disallowed_methods, + reason = "a dev-only installer tool that never talks to providers" +)] + mod agent; mod error; mod install; diff --git a/litellm-rust/crates/tracing/Cargo.toml b/litellm-rust/crates/tracing/Cargo.toml index 41ad20afb3e..914d988d301 100644 --- a/litellm-rust/crates/tracing/Cargo.toml +++ b/litellm-rust/crates/tracing/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +base64.workspace = true fancy-regex.workspace = true percent-encoding.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/tracing/src/lib.rs b/litellm-rust/crates/tracing/src/lib.rs index 47f97d6db27..4c6ec104f3a 100644 --- a/litellm-rust/crates/tracing/src/lib.rs +++ b/litellm-rust/crates/tracing/src/lib.rs @@ -5,6 +5,7 @@ use std::{ pin::pin, }; +use base64::{Engine, engine::general_purpose::STANDARD}; use serde_json::{Map, Value}; use tracing::{ Dispatch, Event, Subscriber, @@ -20,6 +21,31 @@ pub use processing::{DiagnosticInput, DiagnosticOutput, Policy, Processor}; pub use redaction::{REDACTED, SecretRedactor}; pub use tracing::{Level, Metadata, debug, error, info, trace, warn}; +pub struct ByteChunk<'a>(&'a [u8]); + +impl<'a> ByteChunk<'a> { + pub fn new(data: &'a [u8]) -> Self { + Self(data) + } + + pub fn encoding(&self) -> &'static str { + if std::str::from_utf8(self.0).is_ok() { + "utf8" + } else { + "base64" + } + } +} + +impl fmt::Display for ByteChunk<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match std::str::from_utf8(self.0) { + Ok(text) => formatter.write_str(text), + Err(_) => formatter.write_str(&STANDARD.encode(self.0)), + } + } +} + pub trait Sink: Send + Sync + 'static { fn enabled(&self, metadata: &Metadata<'_>) -> bool; fn emit(&self, record: &Record); @@ -44,6 +70,10 @@ impl Logger { } } + pub fn install_global(&self) -> Result<(), tracing::dispatcher::SetGlobalDefaultError> { + tracing::dispatcher::set_global_default(self.dispatch.clone()) + } + pub fn scope(&self, operation: impl FnOnce() -> T) -> T { if EMITTING.get() { return operation(); diff --git a/litellm-rust/crates/tracing/tests/logging.rs b/litellm-rust/crates/tracing/tests/logging.rs index 585e442dad1..3387f259f8b 100644 --- a/litellm-rust/crates/tracing/tests/logging.rs +++ b/litellm-rust/crates/tracing/tests/logging.rs @@ -4,7 +4,9 @@ use std::sync::{ mpsc, }; -use litellm_tracing::{Level, Logger, Metadata, Record, Sink, info, warn}; +use base64::{Engine, engine::general_purpose::STANDARD}; +use litellm_tracing::{ByteChunk, Level, Logger, Metadata, Record, Sink, info, warn}; +use rstest::rstest; use serde_json::{Value, json}; struct Output { @@ -120,3 +122,18 @@ fn nested_scopes_restore_the_previous_sink() { ["inside"] ); } + +#[rstest] +#[case::utf8(b"event: message_stop\n\n", "utf8")] +#[case::binary(&[0xff, 0x00, 0x80], "base64")] +fn byte_chunk_logging_preserves_exact_bytes(#[case] bytes: &[u8], #[case] encoding: &str) { + let chunk = ByteChunk::new(bytes); + assert_eq!(chunk.encoding(), encoding); + let text = chunk.to_string(); + let recovered = match encoding { + "utf8" => text.into_bytes(), + "base64" => STANDARD.decode(text).unwrap(), + _ => unreachable!(), + }; + assert_eq!(recovered, bytes); +} diff --git a/litellm-rust/crates/types/Cargo.toml b/litellm-rust/crates/types/Cargo.toml index 0a0927386f0..e356c8e127d 100644 --- a/litellm-rust/crates/types/Cargo.toml +++ b/litellm-rust/crates/types/Cargo.toml @@ -5,9 +5,14 @@ edition.workspace = true license.workspace = true repository.workspace = true +[features] +schema = ["dep:schemars"] + [dependencies] +schemars = { workspace = true, optional = true } serde.workspace = true serde_json.workspace = true +strum.workspace = true [dev-dependencies] rstest.workspace = true diff --git a/litellm-rust/crates/types/src/lib.rs b/litellm-rust/crates/types/src/lib.rs index da5c9ea893f..dc00ea7128e 100644 --- a/litellm-rust/crates/types/src/lib.rs +++ b/litellm-rust/crates/types/src/lib.rs @@ -1,3 +1,4 @@ pub mod llms; +pub mod recognized; pub mod responses; pub mod utils; diff --git a/litellm-rust/crates/types/src/llms/anthropic.rs b/litellm-rust/crates/types/src/llms/anthropic.rs new file mode 100644 index 00000000000..3e0c4369b11 --- /dev/null +++ b/litellm-rust/crates/types/src/llms/anthropic.rs @@ -0,0 +1,240 @@ +use std::{ + cmp::Ordering, + collections::BTreeSet, + convert::Infallible, + fmt, + hash::{Hash, Hasher}, + str::FromStr, +}; + +/// One value of the `anthropic-beta` header. Equality, ordering and hashing follow the wire +/// string, so a value parsed from a caller's header never disagrees with the matching variant. +#[derive(Clone, Debug, strum::AsRefStr, strum::Display, strum::EnumString)] +pub enum AnthropicBeta { + #[strum(serialize = "oauth-2025-04-20")] + Oauth20250420, + #[strum(serialize = "web-fetch-2025-09-10")] + WebFetch20250910, + #[strum(serialize = "web-search-2025-03-05")] + WebSearch20250305, + #[strum(serialize = "context-management-2025-06-27")] + ContextManagement20250627, + #[strum(serialize = "compact-2026-01-12")] + Compact20260112, + #[strum(serialize = "compact-2026-09-04")] + Compact20260904, + #[strum(serialize = "structured-outputs-2025-11-13")] + StructuredOutputs20251113, + #[strum(serialize = "advanced-tool-use-2025-11-20")] + AdvancedToolUse20251120, + #[strum(serialize = "fast-mode-2026-02-01")] + FastMode20260201, + #[strum(serialize = "advisor-tool-2026-03-01")] + AdvisorTool20260301, + #[strum(serialize = "per-turn-control-2026-07-01")] + PerTurnControl20260701, + #[strum(serialize = "dangerous-tool-use-2026-09-03")] + DangerousToolUse20260903, + #[strum(default, transparent)] + Other(String), +} + +impl AnthropicBeta { + pub const KNOWN: [Self; 12] = [ + Self::Oauth20250420, + Self::WebFetch20250910, + Self::WebSearch20250305, + Self::ContextManagement20250627, + Self::Compact20260112, + Self::Compact20260904, + Self::StructuredOutputs20251113, + Self::AdvancedToolUse20251120, + Self::FastMode20260201, + Self::AdvisorTool20260301, + Self::PerTurnControl20260701, + Self::DangerousToolUse20260903, + ]; + + pub fn as_str(&self) -> &str { + self.as_ref() + } +} + +impl PartialEq for AnthropicBeta { + fn eq(&self, other: &Self) -> bool { + self.as_str() == other.as_str() + } +} + +impl Eq for AnthropicBeta {} + +impl Hash for AnthropicBeta { + fn hash(&self, state: &mut H) { + self.as_str().hash(state); + } +} + +impl PartialOrd for AnthropicBeta { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +impl Ord for AnthropicBeta { + fn cmp(&self, other: &Self) -> Ordering { + self.as_str().cmp(other.as_str()) + } +} + +/// The values of one `anthropic-beta` header: sorted, deduplicated, comma-joined on the wire. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct BetaSet(BTreeSet); + +impl BetaSet { + pub fn is_empty(&self) -> bool { + self.0.is_empty() + } + + pub fn contains(&self, beta: &AnthropicBeta) -> bool { + self.0.contains(beta) + } + + pub fn iter(&self) -> impl Iterator { + self.0.iter() + } + + pub fn union(self, other: Self) -> Self { + self.0.into_iter().chain(other.0).collect() + } +} + +impl FromIterator for BetaSet { + fn from_iter>(betas: I) -> Self { + Self(betas.into_iter().collect()) + } +} + +impl IntoIterator for BetaSet { + type Item = AnthropicBeta; + type IntoIter = std::collections::btree_set::IntoIter; + + fn into_iter(self) -> Self::IntoIter { + self.0.into_iter() + } +} + +impl FromStr for BetaSet { + type Err = Infallible; + + fn from_str(header: &str) -> Result { + Ok(header + .split(',') + .map(str::trim) + .filter(|piece| !piece.is_empty()) + .map(|piece| AnthropicBeta::from_str(piece).unwrap_or_else(|never| match never {})) + .collect()) + } +} + +impl fmt::Display for BetaSet { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut betas = self.0.iter(); + let Some(first) = betas.next() else { + return Ok(()); + }; + f.write_str(first.as_str())?; + betas.try_for_each(|beta| write!(f, ",{beta}")) + } +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + + use super::*; + + fn set(header: &str) -> BetaSet { + header.parse().unwrap_or_else(|never| match never {}) + } + + #[rstest] + fn every_known_beta_parses_back_to_itself( + #[values( + AnthropicBeta::Oauth20250420, + AnthropicBeta::WebFetch20250910, + AnthropicBeta::WebSearch20250305, + AnthropicBeta::ContextManagement20250627, + AnthropicBeta::Compact20260112, + AnthropicBeta::Compact20260904, + AnthropicBeta::StructuredOutputs20251113, + AnthropicBeta::AdvancedToolUse20251120, + AnthropicBeta::FastMode20260201, + AnthropicBeta::AdvisorTool20260301, + AnthropicBeta::PerTurnControl20260701, + AnthropicBeta::DangerousToolUse20260903 + )] + beta: AnthropicBeta, + ) { + let parsed: AnthropicBeta = beta.as_str().parse().unwrap(); + assert!(!matches!(parsed, AnthropicBeta::Other(_))); + assert_eq!(parsed, beta); + assert!(AnthropicBeta::KNOWN.contains(&beta)); + } + + #[test] + fn unknown_values_are_kept_verbatim() { + let parsed: AnthropicBeta = "claude-code-20250219".parse().unwrap(); + assert_eq!( + parsed, + AnthropicBeta::Other("claude-code-20250219".to_string()) + ); + assert_eq!(parsed.to_string(), "claude-code-20250219"); + } + + #[test] + fn a_known_value_spelled_as_other_is_the_same_beta() { + let spelled_out = AnthropicBeta::Other("compact-2026-01-12".to_string()); + assert_eq!(spelled_out, AnthropicBeta::Compact20260112); + assert_eq!( + spelled_out.cmp(&AnthropicBeta::Compact20260112), + Ordering::Equal + ); + assert_eq!( + BetaSet::from_iter([spelled_out, AnthropicBeta::Compact20260112]).to_string(), + "compact-2026-01-12" + ); + } + + #[rstest] + #[case::empty("", "")] + #[case::blank_pieces(" , ,", "")] + #[case::single("b", "b")] + #[case::sorted("c,a", "a,c")] + #[case::trimmed_and_deduplicated("b, a ,b", "a,b")] + #[case::blank_pieces_skipped("a,,b", "a,b")] + #[case::known_and_unknown_sort_together( + "web-search-2025-03-05,claude-code-20250219,fast-mode-2026-02-01", + "claude-code-20250219,fast-mode-2026-02-01,web-search-2025-03-05" + )] + fn header_values_round_trip_sorted_and_deduplicated(#[case] header: &str, #[case] wire: &str) { + assert_eq!(set(header).to_string(), wire); + assert_eq!(set(header).is_empty(), wire.is_empty()); + } + + #[rstest] + #[case::disjoint("a,c", "b", "a,b,c")] + #[case::overlapping("a,b", "b,c", "a,b,c")] + #[case::empty_right("a", "", "a")] + #[case::empty_left("", "a", "a")] + fn union_merges_both_sides(#[case] left: &str, #[case] right: &str, #[case] wire: &str) { + assert_eq!(set(left).union(set(right)).to_string(), wire); + } + + #[test] + fn contains_matches_by_wire_value() { + let betas = set("oauth-2025-04-20,claude-code-20250219"); + assert!(betas.contains(&AnthropicBeta::Oauth20250420)); + assert!(betas.contains(&AnthropicBeta::Other("claude-code-20250219".into()))); + assert!(!betas.contains(&AnthropicBeta::FastMode20260201)); + } +} diff --git a/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs b/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs index 2f7a75ba517..335a03c4b5b 100644 --- a/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs +++ b/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs @@ -1,5 +1,8 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use strum::IntoStaticStr; + +use crate::{llms::openai::ReasoningEffort, recognized::Recognized}; #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[serde(untagged)] @@ -79,10 +82,181 @@ pub struct AnthropicMessage { pub extra: Map, } +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +#[strum(serialize_all = "lowercase")] +pub enum EffortLevel { + Low, + Medium, + High, + Xhigh, + Max, +} + +impl EffortLevel { + pub fn as_str(self) -> &'static str { + self.into() + } +} + +impl From for ReasoningEffort { + fn from(level: EffortLevel) -> Self { + match level { + EffortLevel::Low => Self::Low, + EffortLevel::Medium => Self::Medium, + EffortLevel::High => Self::High, + EffortLevel::Xhigh => Self::Xhigh, + EffortLevel::Max => Self::Max, + } + } +} + +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +#[strum(serialize_all = "lowercase")] +pub enum Speed { + Fast, + Standard, +} + +impl Speed { + pub fn as_str(self) -> &'static str { + self.into() + } +} + +/// The tools whose presence changes how the request is sent. Every other tool, custom or +/// server, deserializes as `Recognized::Unrecognized` and passes through verbatim. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "type")] +pub enum AnthropicTool { + #[serde(rename = "advisor_20260301")] + Advisor { + #[serde(flatten)] + extra: Map, + }, + #[serde(rename = "tool_search_tool_regex_20251119")] + ToolSearchRegex { + #[serde(flatten)] + extra: Map, + }, + #[serde(rename = "tool_search_tool_bm25_20251119")] + ToolSearchBm25 { + #[serde(flatten)] + extra: Map, + }, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "type")] +pub enum ContextEdit { + #[serde(rename = "compact_20260112")] + Compact { + #[serde(flatten)] + extra: Map, + }, + #[serde(rename = "clear_tool_uses_20250919")] + ClearToolUses { + #[serde(flatten)] + extra: Map, + }, + #[serde(rename = "clear_thinking_20251015")] + ClearThinking { + #[serde(flatten)] + extra: Map, + }, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct ContextManagement { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub edits: Option>>, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct OutputConfig { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub effort: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub format: Option, + #[serde(flatten)] + pub extra: Map, +} + +impl OutputConfig { + pub fn is_empty(&self) -> bool { + self.effort.is_none() && self.format.is_none() && self.extra.is_empty() + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum ThinkingDisplay { + Summarized, + Omitted, + Updates, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct EnabledThinking { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub budget_tokens: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub display: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct AdaptiveThinking { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub display: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct DisabledThinking { + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "lowercase")] +pub enum ThinkingConfig { + Enabled(EnabledThinking), + Adaptive(AdaptiveThinking), + Disabled(DisabledThinking), +} + +impl ThinkingConfig { + pub fn enabled(budget_tokens: u64) -> Self { + Self::Enabled(EnabledThinking { + budget_tokens: Some(Recognized::Known(budget_tokens)), + ..EnabledThinking::default() + }) + } + + pub fn adaptive(display: Option) -> Self { + Self::Adaptive(AdaptiveThinking { + display: display.map(Recognized::Known), + ..AdaptiveThinking::default() + }) + } +} + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct AnthropicMessagesRequest { pub model: String, pub messages: Vec, + #[serde(flatten)] + pub params: AnthropicMessagesOptionalParams, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct AnthropicMessagesOptionalParams { #[serde(skip_serializing_if = "Option::is_none")] pub max_tokens: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -100,11 +274,11 @@ pub struct AnthropicMessagesRequest { #[serde(skip_serializing_if = "Option::is_none")] pub top_k: Option, #[serde(skip_serializing_if = "Option::is_none")] - pub tools: Option>, + pub tools: Option>>, #[serde(skip_serializing_if = "Option::is_none")] pub tool_choice: Option, #[serde(skip_serializing_if = "Option::is_none")] - pub thinking: Option, + pub thinking: Option>, #[serde(skip_serializing_if = "Option::is_none")] pub service_tier: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -112,17 +286,17 @@ pub struct AnthropicMessagesRequest { #[serde(skip_serializing_if = "Option::is_none")] pub mcp_servers: Option>, #[serde(skip_serializing_if = "Option::is_none")] - pub context_management: Option, + pub context_management: Option>, #[serde(skip_serializing_if = "Option::is_none")] pub output_format: Option, #[serde(skip_serializing_if = "Option::is_none")] - pub output_config: Option, + pub output_config: Option>, #[serde(skip_serializing_if = "Option::is_none")] - pub speed: Option, + pub speed: Option>, #[serde(skip_serializing_if = "Option::is_none")] pub inference_geo: Option, #[serde(skip_serializing_if = "Option::is_none")] - pub reasoning_effort: Option, + pub reasoning_effort: Option>, #[serde(skip_serializing_if = "Option::is_none")] pub compaction: Option, #[serde(flatten)] @@ -182,6 +356,33 @@ mod tests { assert_eq!(round_trip::(&block), block); } + #[test] + fn request_splits_required_fields_from_optional_params() { + let body = json!({ + "model": "m", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 16, + "stream": true, + "safeguards": [{"type": "dangerous_tool_use"}] + }); + let request: AnthropicMessagesRequest = serde_json::from_value(body.clone()).unwrap(); + + assert_eq!( + ( + request.params.max_tokens, + request.params.stream, + request + .params + .extra + .keys() + .map(String::as_str) + .collect::>(), + ), + (Some(16_u64), Some(true), vec!["safeguards"]) + ); + assert_eq!(serde_json::to_value(request).unwrap(), body); + } + #[test] fn text_constructor_serializes_as_a_text_block() { assert_eq!( @@ -240,7 +441,223 @@ mod tests { "safeguards": [{"type": "dangerous_tool_use", "classifier_context": {"v": 1}}], "metadata": {"user_id": "u"} }))] + #[case::typed_thinking_and_output_config(json!({ + "model": "m", + "messages": [], + "thinking": {"type": "enabled", "budget_tokens": 2048, "display": "omitted", "block_binding": {"prefix_mismatch_behavior": "drop_block"}}, + "output_config": {"effort": "xhigh", "format": {"type": "json_schema", "schema": {}}, "task_budget": {"type": "tokens", "total": 4096}} + }))] + #[case::unrecognized_values_are_kept_verbatim(json!({ + "model": "m", + "messages": [], + "reasoning_effort": "turbo", + "thinking": {"type": "adaptive", "display": "loud"}, + "output_config": {"effort": 5} + }))] + #[case::unrecognized_shapes_are_kept_verbatim(json!({ + "model": "m", + "messages": [], + "reasoning_effort": 3, + "thinking": {"type": "future", "budget_tokens": 1}, + "output_config": "bogus" + }))] + #[case::tools_speed_and_context_management(json!({ + "model": "m", + "messages": [], + "speed": "fast", + "tools": [ + {"name": "get_weather", "input_schema": {"type": "object"}}, + {"type": "custom", "name": "f", "input_schema": {}}, + {"type": "web_search_20250305", "name": "web_search", "max_uses": 3}, + {"type": "advisor_20260301", "name": "advisor", "model": "claude-opus-4-6"}, + {"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"}, + {"type": "tool_search_tool_bm25_20251119"} + ], + "context_management": {"edits": [ + {"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}}, + {"type": "clear_tool_uses_20250919", "keep": {"type": "tool_uses", "value": 3}}, + {"type": "clear_thinking_20251015"}, + {"type": "future_edit"}, + {} + ], "future": true} + }))] + #[case::unrecognized_tools_speed_and_context_management_are_kept_verbatim(json!({ + "model": "m", + "messages": [], + "speed": "turbo", + "tools": ["none", 5], + "context_management": [{"type": "compaction", "compact_threshold": 5}] + }))] fn request_round_trips_unchanged(#[case] request: Value) { assert_eq!(round_trip::(&request), request); } + + #[rstest] + #[case::enabled( + json!({"type": "enabled", "budget_tokens": 2048, "display": "omitted"}), + ThinkingConfig::Enabled(EnabledThinking { + budget_tokens: Some(Recognized::Known(2048)), + display: Some(Recognized::Known(ThinkingDisplay::Omitted)), + extra: Map::new(), + }) + )] + #[case::enabled_without_budget( + json!({"type": "enabled"}), + ThinkingConfig::Enabled(EnabledThinking::default()) + )] + #[case::enabled_with_unrecognized_budget( + json!({"type": "enabled", "budget_tokens": "lots"}), + ThinkingConfig::Enabled(EnabledThinking { + budget_tokens: Some(Recognized::Unrecognized(json!("lots"))), + ..EnabledThinking::default() + }) + )] + #[case::adaptive_with_unrecognized_display( + json!({"type": "adaptive", "display": "loud"}), + ThinkingConfig::Adaptive(AdaptiveThinking { + display: Some(Recognized::Unrecognized(json!("loud"))), + extra: Map::new(), + }) + )] + #[case::disabled_keeps_extra_fields( + json!({"type": "disabled", "future": true}), + ThinkingConfig::Disabled(DisabledThinking { + extra: Map::from_iter([("future".to_string(), json!(true))]), + }) + )] + fn thinking_config_parses_every_documented_type_leniently( + #[case] thinking: Value, + #[case] expected: ThinkingConfig, + ) { + assert_eq!( + serde_json::from_value::(thinking).unwrap(), + expected + ); + } + + #[rstest] + #[case::advisor( + json!({"type": "advisor_20260301", "name": "advisor"}), + Recognized::Known(AnthropicTool::Advisor { extra: Map::from_iter([("name".to_string(), json!("advisor"))]) }) + )] + #[case::regex_tool_search( + json!({"type": "tool_search_tool_regex_20251119"}), + Recognized::Known(AnthropicTool::ToolSearchRegex { extra: Map::new() }) + )] + #[case::bm25_tool_search( + json!({"type": "tool_search_tool_bm25_20251119"}), + Recognized::Known(AnthropicTool::ToolSearchBm25 { extra: Map::new() }) + )] + #[case::custom_tool_without_a_type( + json!({"name": "advisor", "input_schema": {}}), + Recognized::Unrecognized(json!({"name": "advisor", "input_schema": {}})) + )] + #[case::other_server_tool( + json!({"type": "web_search_20250305", "name": "web_search"}), + Recognized::Unrecognized(json!({"type": "web_search_20250305", "name": "web_search"})) + )] + #[case::not_an_object(json!("advisor_20260301"), Recognized::Unrecognized(json!("advisor_20260301")))] + fn tools_are_recognized_by_their_exact_type( + #[case] tool: Value, + #[case] expected: Recognized, + ) { + assert_eq!( + serde_json::from_value::>(tool).unwrap(), + expected + ); + } + + #[rstest] + #[case::compact( + json!({"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1}}), + Recognized::Known(ContextEdit::Compact { + extra: Map::from_iter([("trigger".to_string(), json!({"type": "input_tokens", "value": 1}))]), + }) + )] + #[case::clear_tool_uses( + json!({"type": "clear_tool_uses_20250919"}), + Recognized::Known(ContextEdit::ClearToolUses { extra: Map::new() }) + )] + #[case::clear_thinking( + json!({"type": "clear_thinking_20251015"}), + Recognized::Known(ContextEdit::ClearThinking { extra: Map::new() }) + )] + #[case::unknown_type(json!({"type": "future"}), Recognized::Unrecognized(json!({"type": "future"})))] + #[case::no_type(json!({}), Recognized::Unrecognized(json!({})))] + fn context_edits_are_recognized_by_their_exact_type( + #[case] edit: Value, + #[case] expected: Recognized, + ) { + assert_eq!( + serde_json::from_value::>(edit).unwrap(), + expected + ); + } + + #[rstest] + #[case::edits( + json!({"edits": [{"type": "compact_20260112"}]}), + Recognized::Known(ContextManagement { + edits: Some(vec![Recognized::Known(ContextEdit::Compact { extra: Map::new() })]), + extra: Map::new(), + }) + )] + #[case::object_without_edits( + json!({"future": 1}), + Recognized::Known(ContextManagement { + edits: None, + extra: Map::from_iter([("future".to_string(), json!(1))]), + }) + )] + #[case::openai_list(json!([{"type": "compaction"}]), Recognized::Unrecognized(json!([{"type": "compaction"}])))] + #[case::edits_not_a_list(json!({"edits": 5}), Recognized::Unrecognized(json!({"edits": 5})))] + #[case::scalar(json!("compaction"), Recognized::Unrecognized(json!("compaction")))] + fn context_management_is_known_only_as_an_edits_object( + #[case] value: Value, + #[case] expected: Recognized, + ) { + assert_eq!( + serde_json::from_value::>(value).unwrap(), + expected + ); + } + + #[rstest] + #[case::fast(json!("fast"), Recognized::Known(Speed::Fast))] + #[case::standard(json!("standard"), Recognized::Known(Speed::Standard))] + #[case::unknown(json!("turbo"), Recognized::Unrecognized(json!("turbo")))] + #[case::wrong_case(json!("Fast"), Recognized::Unrecognized(json!("Fast")))] + #[case::not_a_string(json!(1), Recognized::Unrecognized(json!(1)))] + fn speed_is_known_only_as_a_documented_value( + #[case] value: Value, + #[case] expected: Recognized, + ) { + assert_eq!( + serde_json::from_value::>(value).unwrap(), + expected + ); + } + + #[rstest] + fn speed_names_match_the_wire(#[values(Speed::Fast, Speed::Standard)] speed: Speed) { + assert_eq!(serde_json::to_value(speed).unwrap(), json!(speed.as_str())); + } + + #[rstest] + fn effort_level_names_match_the_wire( + #[values( + EffortLevel::Low, + EffortLevel::Medium, + EffortLevel::High, + EffortLevel::Xhigh, + EffortLevel::Max + )] + level: EffortLevel, + ) { + assert_eq!(serde_json::to_value(level).unwrap(), json!(level.as_str())); + assert_eq!( + serde_json::to_value(ReasoningEffort::from(level)).unwrap(), + json!(level.as_str()) + ); + } } diff --git a/litellm-rust/crates/types/src/llms/mod.rs b/litellm-rust/crates/types/src/llms/mod.rs index 09d2207a0ca..19ce0bb77ef 100644 --- a/litellm-rust/crates/types/src/llms/mod.rs +++ b/litellm-rust/crates/types/src/llms/mod.rs @@ -1,2 +1,3 @@ +pub mod anthropic; pub mod anthropic_messages; pub mod openai; diff --git a/litellm-rust/crates/types/src/llms/openai.rs b/litellm-rust/crates/types/src/llms/openai.rs index 232f5b9cc51..ee8c882c40c 100644 --- a/litellm-rust/crates/types/src/llms/openai.rs +++ b/litellm-rust/crates/types/src/llms/openai.rs @@ -1,5 +1,43 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use strum::IntoStaticStr; + +/// Reasoning effort level accepted or applied by the model. +#[derive(Clone, Copy, Debug, Deserialize, Eq, IntoStaticStr, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] +pub enum ReasoningEffort { + None, + Minimal, + Low, + Medium, + High, + Xhigh, + Max, +} + +impl ReasoningEffort { + pub const ALL: [Self; 7] = [ + Self::None, + Self::Minimal, + Self::Low, + Self::Medium, + Self::High, + Self::Xhigh, + Self::Max, + ]; + + pub fn as_str(self) -> &'static str { + self.into() + } + + pub fn parse(value: &str) -> Option { + Self::ALL + .into_iter() + .find(|effort| effort.as_str() == value) + } +} #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[serde(untagged)] @@ -56,3 +94,39 @@ pub enum ChatCompletionThinkingBlock { cache_control: Option, }, } + +#[cfg(test)] +mod tests { + use rstest::rstest; + + use super::*; + + #[rstest] + fn reasoning_effort_names_match_the_wire_and_parse_back( + #[values( + ReasoningEffort::None, + ReasoningEffort::Minimal, + ReasoningEffort::Low, + ReasoningEffort::Medium, + ReasoningEffort::High, + ReasoningEffort::Xhigh, + ReasoningEffort::Max + )] + effort: ReasoningEffort, + ) { + assert_eq!( + serde_json::to_value(effort).unwrap(), + Value::String(effort.as_str().to_string()) + ); + assert_eq!(ReasoningEffort::parse(effort.as_str()), Some(effort)); + assert!(ReasoningEffort::ALL.contains(&effort)); + } + + #[rstest] + #[case::unknown("ultra")] + #[case::uppercase("HIGH")] + #[case::empty("")] + fn reasoning_effort_parse_rejects(#[case] value: &str) { + assert_eq!(ReasoningEffort::parse(value), None); + } +} diff --git a/litellm-rust/crates/types/src/recognized.rs b/litellm-rust/crates/types/src/recognized.rs new file mode 100644 index 00000000000..d82b51f9fde --- /dev/null +++ b/litellm-rust/crates/types/src/recognized.rs @@ -0,0 +1,42 @@ +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(untagged)] +pub enum Recognized { + Known(T), + Unrecognized(Value), +} + +impl Recognized { + pub fn known(&self) -> Option<&T> { + match self { + Self::Known(value) => Some(value), + Self::Unrecognized(_) => None, + } + } +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + use serde_json::json; + + use super::*; + + #[rstest] + #[case::known(json!(7), Recognized::Known(7))] + #[case::wrong_type(json!("7"), Recognized::Unrecognized(json!("7")))] + #[case::out_of_range(json!(-1), Recognized::Unrecognized(json!(-1)))] + #[case::object(json!({"a": 1}), Recognized::Unrecognized(json!({"a": 1})))] + fn value_is_known_only_when_it_parses_as_the_type( + #[case] value: Value, + #[case] expected: Recognized, + ) { + assert_eq!( + serde_json::from_value::>(value.clone()).unwrap(), + expected + ); + assert_eq!(serde_json::to_value(expected).unwrap(), value); + } +} diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 246ac4fd369..9974e77d017 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -43,7 +43,7 @@ def _uses_native_vertex_output( ) -> bool: if custom_llm_provider != "vertex_ai": return False - if model_name and getattr(litellm, "disable_vertex_batch_output_transformation", False): + if model_name and litellm.disable_vertex_batch_output_transformation: return True return first_row is not None and is_native_vertex_batch_output_row(first_row) diff --git a/litellm/containers/README.md b/litellm/containers/README.md index b54f96b1132..571bcaf415c 100644 --- a/litellm/containers/README.md +++ b/litellm/containers/README.md @@ -213,7 +213,7 @@ Run the container API tests: ```bash cd /Users/ishaanjaffer/github/litellm -python -m pytest tests/test_litellm/containers/ -v +python -m pytest tests/unit/containers/ -v ``` Test via proxy: diff --git a/litellm/integrations/SlackAlerting/budget_alert_types.py b/litellm/integrations/SlackAlerting/budget_alert_types.py index f35ff7b5f82..4fe833acecc 100644 --- a/litellm/integrations/SlackAlerting/budget_alert_types.py +++ b/litellm/integrations/SlackAlerting/budget_alert_types.py @@ -63,6 +63,8 @@ class TokenBudgetAlert(BaseBudgetAlertType): return "Key Budget: " def get_id(self, user_info: CallInfo) -> str: + if user_info.event_group == Litellm_EntityType.TEAM_MEMBER: + return f"team_member:{user_info.user_id}:{user_info.team_id}" return user_info.token or "default_id" diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 17ec3ed787d..7c608aac8d9 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -555,7 +555,11 @@ class SlackAlerting(CustomBatchLogger): budget_alert_class: Final = get_budget_alert_type(type) _id: Final = budget_alert_class.get_id(user_info) user_info_str: Final = self._get_user_info_str(user_info) - event_message = budget_alert_class.get_event_message() + event_message = ( + "Team Member Budget: " + if user_info.event_group == Litellm_EntityType.TEAM_MEMBER + else budget_alert_class.get_event_message() + ) # Set default event unless we're in projected_limit_exceeded event: ( diff --git a/litellm/integrations/email_templates/templates.py b/litellm/integrations/email_templates/templates.py index 935067c97fc..2bd079ef15d 100644 --- a/litellm/integrations/email_templates/templates.py +++ b/litellm/integrations/email_templates/templates.py @@ -131,3 +131,25 @@ MAX_BUDGET_ALERT_EMAIL_TEMPLATE: Final = """ {email_footer} """ + +TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE: Final = """ + LiteLLM Logo + +

Hi,
+ + Team member {member} has reached {percentage}% of their team member budget in team {team_alias}.

+ + Current Spend: {spend}
+ Team Member Budget: {max_budget}
+ Alert Threshold: {alert_threshold} ({percentage}%)
+ +

+ Warning: Once this member reaches their team member budget of {max_budget}, their requests in this team will be rejected. +

+ + You can view usage and manage team member budgets in the LiteLLM Dashboard.

+ + If you have any questions, please send an email to {email_support_contact}

+ + {email_footer} +""" diff --git a/litellm/integrations/langfuse/langfuse_sdk.py b/litellm/integrations/langfuse/langfuse_sdk.py index 66819c95ebf..986f35297d2 100644 --- a/litellm/integrations/langfuse/langfuse_sdk.py +++ b/litellm/integrations/langfuse/langfuse_sdk.py @@ -684,7 +684,7 @@ class LangfuseSpanExporter(SpanExporter): def _round(self, halving: _Halving) -> _Halving: sent: Final = tuple((batch, self._send_batch(batch)) for batch in halving.pending) return _Halving( - pending=tuple(part for batch, outcome in sent if outcome == "too_large" for part in _smaller(batch)), + pending=tuple(chain.from_iterable(_smaller(batch) for batch, outcome in sent if outcome == "too_large")), settled=halving.settled + tuple( SpanExportResult.SUCCESS if outcome == "delivered" else SpanExportResult.FAILURE diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 28d72702f3e..e955c0157c6 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -384,6 +384,7 @@ _DEPLOYMENT_PRICING_KEYS: Final = ( "cache_read_input_token_cost_above_200k_tokens_batches", "cache_read_input_token_cost_above_272k_tokens_batches", "cache_creation_input_token_cost_batches", + "cache_creation_input_token_cost_above_200k_tokens_batches", "cache_creation_input_token_cost_above_272k_tokens_batches", "ocr_cost_per_page", "ocr_cost_per_page_batches", diff --git a/litellm/litellm_core_utils/logging_worker.py b/litellm/litellm_core_utils/logging_worker.py index 2f8e7bdccea..03420b84c22 100644 --- a/litellm/litellm_core_utils/logging_worker.py +++ b/litellm/litellm_core_utils/logging_worker.py @@ -165,7 +165,9 @@ class LoggingWorker: if self._worker_task is None or self._worker_task.done(): self._worker_task = asyncio.create_task(self._worker_loop()) - async def _process_log_task(self, task: LoggingTask, sem: asyncio.Semaphore): + async def _process_log_task( + self, task: LoggingTask, sem: asyncio.Semaphore, queue: "asyncio.Queue[LoggingTask]" + ) -> None: """Runs the logging task and handles cleanup. Releases semaphore when done.""" try: if self._queue is not None: @@ -182,7 +184,7 @@ class LoggingWorker: verbose_logger.exception("LoggingWorker error: %s", e) finally: self._untrack_dequeued(task) - self._queue.task_done() + queue.task_done() finally: # Always release semaphore, even if queue is None sem.release() @@ -219,7 +221,8 @@ class LoggingWorker: async def _worker_loop(self) -> None: """Main worker loop that gets tasks and schedules them to run concurrently.""" try: - if self._queue is None or self._sem is None: + queue: Final = self._queue + if queue is None or self._sem is None: return while True: @@ -227,10 +230,10 @@ class LoggingWorker: # unbounded growth of waiting tasks await self._sem.acquire() try: - task = await self._queue.get() + task = await queue.get() self._track_dequeued(task) # Track each spawned coroutine so we can cancel on shutdown. - processing_task = asyncio.create_task(self._process_log_task(task, self._sem)) + processing_task = asyncio.create_task(self._process_log_task(task, self._sem, queue)) self._running_tasks.add(processing_task) processing_task.add_done_callback(self._running_tasks.discard) except Exception: @@ -497,7 +500,8 @@ class LoggingWorker: """ Clear the queue with a maximum time limit. """ - if self._queue is None: + queue: Final = self._queue + if queue is None: return start_time: Final = asyncio.get_event_loop().time() @@ -509,7 +513,7 @@ class LoggingWorker: break try: - task = self._queue.get_nowait() + task = queue.get_nowait() # Await the coroutine to properly execute and avoid "never awaited" warnings try: await asyncio.wait_for( @@ -522,7 +526,7 @@ class LoggingWorker: finally: # Clear reference to prevent memory leaks task = None - self._queue.task_done() # If you're using join() elsewhere + queue.task_done() except asyncio.QueueEmpty: break diff --git a/litellm/llms/openai/organization_costs.py b/litellm/llms/openai/organization_costs.py index e7fb22f9b19..856072ddb99 100644 --- a/litellm/llms/openai/organization_costs.py +++ b/litellm/llms/openai/organization_costs.py @@ -3,6 +3,7 @@ from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass from datetime import date, datetime, timedelta, timezone +from itertools import chain from types import MappingProxyType from typing import Final, Literal, TypeAlias @@ -126,7 +127,8 @@ async def fetch_openai_daily_costs( return MappingProxyType( { day: sum( - result.amount.value for bucket in buckets if _bucket_day(bucket) == day for result in bucket.results + result.amount.value + for result in chain.from_iterable(bucket.results for bucket in buckets if _bucket_day(bucket) == day) ) for day in days } diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 792e7316c67..4b88072b479 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -1279,8 +1279,11 @@ "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "anthropic.claude-mythos-preview": { - "input_cost_per_token": 0, - "output_cost_per_token": 0, + "cache_creation_input_token_cost": 3.4375e-05, + "cache_creation_input_token_cost_above_1hr": 5.5e-05, + "cache_read_input_token_cost": 2.75e-06, + "input_cost_per_token": 2.75e-05, + "output_cost_per_token": 0.0001375, "litellm_provider": "bedrock", "max_input_tokens": 1000000, "max_output_tokens": 128000, @@ -1289,10 +1292,11 @@ "thinking_always_on": true, "supports_function_calling": true, "supports_vision": true, - "supports_prompt_caching": false, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_tool_choice": true, - "supports_output_config": true + "supports_output_config": true, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "global.anthropic.claude-opus-4-7": { "bedrock_converse_supports_strict_tools": false, @@ -7773,27 +7777,27 @@ "output_cost_per_token_above_272k_tokens_batches": 0.000135 }, "azure/gpt-5.6": { - "cache_creation_input_token_cost": 6.25e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.25e-05, - "cache_creation_input_token_cost_priority": 1.25e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.5e-05, - "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1e-06, - "cache_read_input_token_cost_priority": 1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2e-06, - "input_cost_per_token": 5e-06, - "input_cost_per_token_above_272k_tokens": 1e-05, - "input_cost_per_token_priority": 1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2e-05, + "cache_creation_input_token_cost": 5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1e-05, + "cache_creation_input_token_cost_priority": 1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2e-05, + "cache_read_input_token_cost": 4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8e-07, + "cache_read_input_token_cost_priority": 8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.6e-06, + "input_cost_per_token": 4e-06, + "input_cost_per_token_above_272k_tokens": 8e-06, + "input_cost_per_token_priority": 8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.6e-05, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3e-05, - "output_cost_per_token_above_272k_tokens": 4.5e-05, - "output_cost_per_token_priority": 6e-05, - "output_cost_per_token_above_272k_tokens_priority": 9e-05, + "output_cost_per_token": 2e-05, + "output_cost_per_token_above_272k_tokens": 3e-05, + "output_cost_per_token_priority": 4e-05, + "output_cost_per_token_above_272k_tokens_priority": 6e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -8579,27 +8583,27 @@ "supports_web_search": true }, "azure/us/gpt-5.6": { - "cache_creation_input_token_cost": 6.875e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.375e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.75e-05, - "cache_creation_input_token_cost_priority": 1.375e-05, - "cache_read_input_token_cost": 5.5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1.1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2.2e-06, - "cache_read_input_token_cost_priority": 1.1e-06, - "input_cost_per_token": 5.5e-06, - "input_cost_per_token_above_272k_tokens": 1.1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2.2e-05, - "input_cost_per_token_priority": 1.1e-05, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1.1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2.2e-05, + "cache_creation_input_token_cost_priority": 1.1e-05, + "cache_read_input_token_cost": 4.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8.8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.76e-06, + "cache_read_input_token_cost_priority": 8.8e-07, + "input_cost_per_token": 4.4e-06, + "input_cost_per_token_above_272k_tokens": 8.8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.76e-05, + "input_cost_per_token_priority": 8.8e-06, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3.3e-05, - "output_cost_per_token_above_272k_tokens": 4.95e-05, - "output_cost_per_token_above_272k_tokens_priority": 9.9e-05, - "output_cost_per_token_priority": 6.6e-05, + "output_cost_per_token": 2.2e-05, + "output_cost_per_token_above_272k_tokens": 3.3e-05, + "output_cost_per_token_above_272k_tokens_priority": 6.6e-05, + "output_cost_per_token_priority": 4.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -8985,27 +8989,27 @@ "supports_web_search": true }, "azure/eu/gpt-5.6": { - "cache_creation_input_token_cost": 6.875e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.375e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.75e-05, - "cache_creation_input_token_cost_priority": 1.375e-05, - "cache_read_input_token_cost": 5.5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1.1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2.2e-06, - "cache_read_input_token_cost_priority": 1.1e-06, - "input_cost_per_token": 5.5e-06, - "input_cost_per_token_above_272k_tokens": 1.1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2.2e-05, - "input_cost_per_token_priority": 1.1e-05, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1.1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2.2e-05, + "cache_creation_input_token_cost_priority": 1.1e-05, + "cache_read_input_token_cost": 4.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8.8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.76e-06, + "cache_read_input_token_cost_priority": 8.8e-07, + "input_cost_per_token": 4.4e-06, + "input_cost_per_token_above_272k_tokens": 8.8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.76e-05, + "input_cost_per_token_priority": 8.8e-06, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3.3e-05, - "output_cost_per_token_above_272k_tokens": 4.95e-05, - "output_cost_per_token_above_272k_tokens_priority": 9.9e-05, - "output_cost_per_token_priority": 6.6e-05, + "output_cost_per_token": 2.2e-05, + "output_cost_per_token_above_272k_tokens": 3.3e-05, + "output_cost_per_token_above_272k_tokens_priority": 6.6e-05, + "output_cost_per_token_priority": 4.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -11709,8 +11713,8 @@ "input_cost_per_token": 1.75e-06, "litellm_provider": "azure_ai", "mode": "image_generation", - "output_cost_per_image": 0.0338, - "output_cost_per_image_token": 3.3e-05, + "output_cost_per_image": 0.02, + "output_cost_per_image_token": 1.95e-05, "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", "supported_endpoints": [ "/v1/images/generations", @@ -14640,14 +14644,18 @@ "claude-haiku-4-5-20251001": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_creation_input_token_cost_batches": 6.25e-07, "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 5e-06, + "output_cost_per_token_batches": 2.5e-06, "supports_assistant_prefill": true, "supports_function_calling": true, "supports_native_structured_output": true, @@ -14663,14 +14671,18 @@ "claude-haiku-4-5": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_creation_input_token_cost_batches": 6.25e-07, "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 5e-06, + "output_cost_per_token_batches": 2.5e-06, "supports_assistant_prefill": true, "supports_function_calling": true, "supports_native_structured_output": true, @@ -14693,13 +14705,21 @@ "input_cost_per_token_above_200k_tokens": 6e-06, "output_cost_per_token_above_200k_tokens": 2.25e-05, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_creation_input_token_cost_above_200k_tokens_batches": 3.75e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost_above_200k_tokens": 6e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, + "input_cost_per_token_above_200k_tokens_batches": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_above_200k_tokens_batches": 1.125e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14727,13 +14747,21 @@ "input_cost_per_token_above_200k_tokens": 6e-06, "output_cost_per_token_above_200k_tokens": 2.25e-05, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_creation_input_token_cost_above_200k_tokens_batches": 3.75e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost_above_200k_tokens": 6e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, + "input_cost_per_token_above_200k_tokens_batches": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_above_200k_tokens_batches": 1.125e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14757,14 +14785,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, + "cache_creation_input_token_cost_batches": 1.25e-06, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14797,14 +14829,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 3.75e-06, "cache_creation_input_token_cost_above_1hr": 6e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, "input_cost_per_token": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14865,14 +14901,18 @@ "claude-opus-4-5-20251101": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14895,14 +14935,18 @@ "claude-opus-4-5": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14927,14 +14971,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14966,14 +15014,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15004,14 +15056,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15044,14 +15100,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15083,14 +15143,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 1e-06, + "cache_read_input_token_cost_batches": 5e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15123,14 +15187,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_batches": 1.25e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15164,14 +15232,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 5e-06, "cache_creation_input_token_cost_above_1hr": 8e-06, + "cache_creation_input_token_cost_batches": 2.5e-06, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 4e-06, + "input_cost_per_token_batches": 2e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2e-05, + "output_cost_per_token_batches": 1e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15207,14 +15279,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15250,14 +15326,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15756,6 +15836,16 @@ "supports_function_calling": true, "supports_tool_choice": true }, + "c4ai-aya-expanse-32b": { + "input_cost_per_token": 5e-07, + "output_cost_per_token": 1.5e-06, + "litellm_provider": "cohere_chat", + "max_input_tokens": 128000, + "max_output_tokens": 4000, + "max_tokens": 4000, + "mode": "chat", + "source": "https://docs.cohere.com/docs/models" + }, "command-a-plus-05-2026": { "input_cost_per_token": 0.0, "litellm_provider": "cohere_chat", @@ -19104,6 +19194,41 @@ "supports_tool_choice": true, "supports_vision": true }, + "databricks/databricks-claude-opus-5-5": { + "cache_creation_input_token_cost": 5.00003e-06, + "cache_creation_input_token_cost_above_1hr": 8.00002e-06, + "cache_read_input_token_cost": 1.9999e-07, + "input_cost_per_token": 4.00001e-06, + "input_dbu_cost_per_token": 5.7143e-05, + "litellm_provider": "databricks", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Costs per token are the published Global DBU rates times $0.070 per DBU. The '*_dbu_cost_per_token' fields are provided for reference; cost calculation reads the dollar '*_cost_per_token' fields." + }, + "mode": "chat", + "output_cost_per_token": 1.999998e-05, + "output_dbu_cost_per_token": 0.000285714, + "prompt_cache_min_tokens": 512, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_adaptive_thinking": true, + "supports_anthropic_thinking_payload": true, + "supports_assistant_prefill": false, + "supports_forced_tool_use": false, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_output_config": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "thinking_always_on": true + }, "databricks/databricks-claude-sonnet-4": { "cache_creation_input_token_cost": 3.74997e-06, "cache_read_input_token_cost": 3.0002e-07, @@ -22145,11 +22270,11 @@ } }, "bing_grounding/search": { - "input_cost_per_query": 0.035, + "input_cost_per_query": 0.014, "litellm_provider": "bing_grounding", "mode": "search", "metadata": { - "notes": "Grounding with Bing Search (G1 SKU): $35 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment." + "notes": "Grounding with Bing Search (G1 SKU): $14 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment." } }, "tinyfish/search": { @@ -28036,6 +28161,7 @@ "tpm": 10000000 }, "gemini/gemini-3-pro-image-preview": { + "deprecation_date": "2026-06-25", "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, @@ -28083,6 +28209,7 @@ "supports_reasoning": false }, "gemini/gemini-3.1-flash-image-preview": { + "deprecation_date": "2026-06-25", "input_cost_per_token": 5e-07, "input_cost_per_token_batches": 2.5e-07, "litellm_provider": "gemini", @@ -28130,6 +28257,7 @@ "cache_read_input_token_cost_batches": 1.25e-08, "cache_read_input_token_cost_flex": 1.25e-08, "cache_read_input_token_cost_priority": 4.5e-08, + "deprecation_date": "2026-05-25", "input_cost_per_audio_token": 5e-07, "input_cost_per_token": 2.5e-07, "input_cost_per_token_batches": 1.25e-07, @@ -30551,7 +30679,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.6-luna": { "litellm_provider": "chatgpt", @@ -30567,7 +30699,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.6-sol": { "litellm_provider": "chatgpt", @@ -30583,7 +30719,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.6-terra": { "litellm_provider": "chatgpt", @@ -30599,7 +30739,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.4": { "litellm_provider": "chatgpt", @@ -30614,7 +30758,12 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "default_reasoning_effort": "none", + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.4-pro": { "litellm_provider": "chatgpt", @@ -30628,7 +30777,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": false, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.3-codex": { "litellm_provider": "chatgpt", @@ -30642,7 +30795,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": false, + "supports_xhigh_reasoning_effort": false, + "supports_minimal_reasoning_effort": true }, "chatgpt/gpt-5.3-codex-spark": { "litellm_provider": "chatgpt", @@ -30686,7 +30843,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.2-codex": { "litellm_provider": "chatgpt", @@ -30700,7 +30858,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.2": { "litellm_provider": "chatgpt", @@ -30715,7 +30874,12 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "default_reasoning_effort": "none", + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.1-codex-max": { "litellm_provider": "chatgpt", @@ -30729,7 +30893,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.1-codex-mini": { "litellm_provider": "chatgpt", @@ -30743,7 +30908,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "gigachat/GigaChat-2": { "input_cost_per_token": 0.0, @@ -39214,6 +39380,17 @@ "supports_function_calling": true, "supports_reasoning": true }, + "nebius/deepseek-ai/DeepSeek-V4.1-Flash": { + "input_cost_per_token": 3e-07, + "litellm_provider": "nebius", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://tokenfactory.nebius.com/endpoints?modals=endpoint-details&model-id=deepseek-ai/DeepSeek-V4.1-Flash", + "supports_vision": true + }, "nebius/MiniMaxAI/MiniMax-M2.5": { "max_tokens": 196608, "max_input_tokens": 196608, @@ -41477,65 +41654,63 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 8.44944e-07, + "cache_read_input_token_cost": 3.828e-08, + "input_cost_per_token": 4.5936e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.689888e-06, + "output_cost_per_token": 9.1872e-07, "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, "supports_function_calling": true, + "supports_pdf_input": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 7.0412e-08, - "supports_audio_input": false, - "supports_pdf_input": false, "supports_vision": false, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "input_cost_per_token": 3e-07, - "output_cost_per_token": 1.2e-06, - "cache_read_input_token_cost": 6e-09, + "cache_read_input_token_cost": 4.2e-09, + "input_cost_per_token": 1.4e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"00:00-01:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"04:00-06:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"10:00-00:00"}],"input_cost_per_token":1.5e-7,"output_cost_per_token":6e-7,"cache_read_input_token_cost":3e-9}, + "output_cost_per_token": 4.2e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_vision": true, "supports_pdf_input": false, "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "input_cost_per_token": 4.62e-07, + "cache_read_input_token_cost": 8.8e-09, + "input_cost_per_token": 2.64e-07, "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.386e-06, + "output_cost_per_token": 7.92e-07, "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, "supports_function_calling": true, + "supports_pdf_input": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 1.54e-08, - "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"00:00-01:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"04:00-06:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"10:00-00:00"}],"input_cost_per_token":0.00000132,"output_cost_per_token":0.00000396,"cache_read_input_token_cost":4.4e-8}, - "supports_audio_input": false, - "supports_pdf_input": false, "supports_vision": false, "supports_web_search": false }, @@ -42755,14 +42930,14 @@ "openrouter/qwen/qwen3-coder-plus": { "cache_creation_input_token_cost": 8.125e-07, "cache_creation_input_token_cost_above_128k_tokens": 2.4375e-06, - "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, - "input_cost_per_token_above_32k_tokens": 1.17e-06, "cache_creation_input_token_cost_above_32k_tokens": 1.4625e-06, - "cache_read_input_token_cost_above_32k_tokens": 2.34e-07, - "output_cost_per_token_above_32k_tokens": 5.85e-06, "cache_read_input_token_cost": 1.3e-07, + "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, + "cache_read_input_token_cost_above_32k_tokens": 2.34e-07, + "deprecation_date": "2026-10-09", "input_cost_per_token": 6.5e-07, "input_cost_per_token_above_128k_tokens": 1.95e-06, + "input_cost_per_token_above_32k_tokens": 1.17e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, "max_output_tokens": 65536, @@ -42770,6 +42945,7 @@ "mode": "chat", "output_cost_per_token": 3.25e-06, "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 5.85e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42802,6 +42978,7 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-235b-a22b-thinking-2507": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2.3e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, @@ -43099,25 +43276,25 @@ "supports_web_search": false }, "openrouter/z-ai/glm-4.7": { - "input_cost_per_token": 4e-07, - "output_cost_per_token": 1.75e-06, "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 8e-08, + "cache_read_input_token_cost": 1.1e-07, + "input_cost_per_token": 6e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", + "output_cost_per_token": 2.2e-06, "source": "https://openrouter.ai/api/v1/models", - "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_vision": false, - "supports_prompt_caching": true, "supports_assistant_prefill": true, "supports_audio_input": false, + "supports_function_calling": true, "supports_pdf_input": false, + "supports_prompt_caching": true, + "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/z-ai/glm-4.7-flash": { @@ -43162,15 +43339,15 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.1": { - "input_cost_per_token": 9.66e-07, - "output_cost_per_token": 3.036e-06, - "cache_read_input_token_cost": 1.794e-07, "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 1.7914e-07, + "input_cost_per_token": 9.646e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", + "output_cost_per_token": 3.0316e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -56230,7 +56407,7 @@ "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", - "max_input_tokens": 1048576, + "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "realtime", @@ -56417,7 +56594,7 @@ "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", - "max_input_tokens": 1048576, + "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "realtime", @@ -59653,14 +59830,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 1e-06, + "cache_read_input_token_cost_batches": 5e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "prompt_cache_min_tokens": 512, "search_context_cost_per_query": { "search_context_size_high": 0.01, @@ -59693,14 +59874,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_batches": 1.25e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -60254,7 +60439,7 @@ "mode": "chat", "output_cost_per_token": 6.6e-07, "output_cost_per_token_priority": 8.25e-07, - "source": "https://api.fireworks.ai/v1/serverless/models", + "source": "https://api.fireworks.ai/v1/serverless/models?format=nested", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -60284,13 +60469,16 @@ }, "fireworks_ai/accounts/fireworks/models/deepseek-v4-flash-vision-exp": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "deprecation_date": "2026-09-25", "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, + "output_cost_per_token_priority": 8.25e-07, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, @@ -60353,7 +60541,7 @@ "mode": "chat", "output_cost_per_token": 6.6e-07, "output_cost_per_token_priority": 8.25e-07, - "source": "https://api.fireworks.ai/v1/serverless/models", + "source": "https://api.fireworks.ai/v1/serverless/models?format=nested", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -60383,13 +60571,16 @@ }, "fireworks_ai/deepseek-v4-flash-vision-exp": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "deprecation_date": "2026-09-25", "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, + "output_cost_per_token_priority": 8.25e-07, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, @@ -60518,14 +60709,17 @@ }, "fireworks_ai/muse-glimmer-30b": { "cache_read_input_token_cost": 4e-08, + "cache_read_input_token_cost_priority": 6e-08, "deprecation_date": "2026-09-25", "input_cost_per_token": 3.5e-07, + "input_cost_per_token_priority": 5.25e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, @@ -60567,14 +60761,17 @@ }, "fireworks_ai/accounts/fireworks/models/muse-glimmer-30b": { "cache_read_input_token_cost": 4e-08, + "cache_read_input_token_cost_priority": 6e-08, "deprecation_date": "2026-09-25", "input_cost_per_token": 3.5e-07, + "input_cost_per_token_priority": 5.25e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, @@ -62988,6 +63185,22 @@ "image" ] }, + "xai/grok-imagine-image-pro": { + "input_cost_per_image": 0.05, + "litellm_provider": "xai", + "mode": "image_generation", + "source": "https://docs.x.ai/docs/models", + "supported_endpoints": [ + "/v1/images/generations" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "image" + ] + }, "xai/grok-imagine-image-2.0": { "input_cost_per_image": 0.06, "litellm_provider": "xai", @@ -63994,6 +64207,62 @@ "supports_response_schema": true, "supports_vision": true }, + "azure_ai/deepseek-r1": { + "deprecation_date": "2026-08-13", + "input_cost_per_token": 1.35e-06, + "output_cost_per_token": 5.4e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/deepseek-v3-0324": { + "deprecation_date": "2026-07-13", + "input_cost_per_token": 1.14e-06, + "output_cost_per_token": 4.56e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/deepseek-v3.1": { + "deprecation_date": "2026-07-13", + "input_cost_per_token": 1.23e-06, + "output_cost_per_token": 4.94e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-3": { + "deprecation_date": "2026-05-01", + "input_cost_per_token": 3e-06, + "output_cost_per_token": 1.5e-05, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-3-mini": { + "deprecation_date": "2026-05-01", + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.27e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-4-fast-non-reasoning": { + "deprecation_date": "2026-05-01", + "input_cost_per_token": 2e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-4-fast-reasoning": { + "deprecation_date": "2026-05-01", + "input_cost_per_token": 2e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, "bedrock/us-gov-west-1/nvidia.nemotron-nano-3-30b": { "input_cost_per_token": 7.2e-08, "litellm_provider": "bedrock", @@ -64822,6 +65091,131 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "bedrock_mantle/deepseek.v3.1": { + "input_cost_per_token": 5.8e-07, + "output_cost_per_token": 1.68e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 128000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-1.html" + }, + "bedrock_mantle/moonshotai.kimi-k2-thinking": { + "input_cost_per_token": 6e-07, + "output_cost_per_token": 2.5e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html" + }, + "bedrock_mantle/qwen.qwen3-235b-a22b-2507": { + "input_cost_per_token": 2.2e-07, + "output_cost_per_token": 8.8e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-235b-a22b-2507.html" + }, + "bedrock_mantle/qwen.qwen3-32b": { + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 32000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-32b.html" + }, + "bedrock_mantle/qwen.qwen3-coder-30b-a3b-instruct": { + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-coder-30b-a3b-instruct.html" + }, + "bedrock_mantle/qwen.qwen3-coder-480b-a35b-instruct": { + "input_cost_per_token": 4.5e-07, + "output_cost_per_token": 1.8e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 128000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-coder-480b-a35b-instruct.html" + }, + "bedrock_mantle/qwen.qwen3-next-80b-a3b-instruct": { + "input_cost_per_token": 1.4e-07, + "output_cost_per_token": 1.2e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-next-80b-a3b.html" + }, + "bedrock_mantle/qwen.qwen3-vl-235b-a22b-instruct": { + "input_cost_per_token": 5.3e-07, + "output_cost_per_token": 2.66e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-vl-235b-a22b.html" + }, "azure/us-gov/gpt-5.1": { "cache_read_input_token_cost": 1.71875e-07, "default_reasoning_effort": "none", @@ -66172,23 +66566,23 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3-flash": { + "cache_read_input_token_cost": 1e-08, "input_cost_per_token": 4.5e-08, - "output_cost_per_token": 6e-07, - "cache_read_input_token_cost": 2.85e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", + "output_cost_per_token": 1.4e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-flash-vision-exp": { @@ -66349,24 +66743,24 @@ "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-flash-0731": { - "input_cost_per_token": 3e-08, - "output_cost_per_token": 3.2e-07, "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 2.2e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", + "output_cost_per_token": 3.2e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_response_schema": true, "supports_parallel_function_calling": true, "supports_pdf_input": false, - "supports_vision": false, "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/qwen/qwen3.7-flash": { @@ -66839,46 +67233,47 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.6-max-preview": { - "input_cost_per_token": 1.027e-06, - "output_cost_per_token": 6.162e-06, "cache_creation_input_token_cost": 1.28375e-06, - "input_cost_per_token_above_128k_tokens": 1.58e-06, - "output_cost_per_token_above_128k_tokens": 9.48e-06, "cache_creation_input_token_cost_above_128k_tokens": 1.975e-06, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 1.027e-06, + "input_cost_per_token_above_128k_tokens": 1.58e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 6.162e-06, + "output_cost_per_token_above_128k_tokens": 9.48e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, "openrouter/qwen/qwen3.6-27b": { - "input_cost_per_token": 3.2e-07, - "output_cost_per_token": 2.7e-06, "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_token": 3.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 262140, "max_tokens": 262140, "mode": "chat", + "output_cost_per_token": 3.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/openai/gpt-5.5-pro": { @@ -66923,23 +67318,23 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "input_cost_per_token": 4.9e-08, - "output_cost_per_token": 9.8e-08, - "cache_read_input_token_cost": 9.8e-09, + "cache_read_input_token_cost": 9.408e-09, + "input_cost_per_token": 4.704e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", + "output_cost_per_token": 9.408e-08, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/moonshotai/kimi-k2.6": { @@ -66964,22 +67359,22 @@ "supports_web_search": false }, "openrouter/google/gemma-4-26b-a4b-it": { - "cache_read_input_token_cost": 5e-08, - "input_cost_per_token": 9e-08, - "output_cost_per_token": 3e-07, + "cache_read_input_token_cost": 3.75e-08, + "input_cost_per_token": 6.75e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, "max_tokens": 235929, "mode": "chat", + "output_cost_per_token": 2.25e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": true, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67264,25 +67659,26 @@ "supports_video_input": true }, "openrouter/qwen/qwen3-max-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 7.8e-07, - "input_cost_per_token_above_32k_tokens": 1.56e-06, - "output_cost_per_token_above_32k_tokens": 7.8e-06, - "output_cost_per_token": 3.9e-06, "input_cost_per_token_above_128k_tokens": 1.95e-06, - "output_cost_per_token_above_128k_tokens": 9.75e-06, + "input_cost_per_token_above_32k_tokens": 1.56e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 3.9e-06, + "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 7.8e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -67534,59 +67930,62 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-vl-32b-instruct": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.04e-07, - "output_cost_per_token": 4.16e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4.16e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-8b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.8e-07, - "output_cost_per_token": 2.1e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.1e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-8b-instruct": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.17e-07, - "output_cost_per_token": 4.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4.55e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67616,40 +68015,41 @@ "supports_web_search": true }, "openrouter/qwen/qwen3-vl-30b-a3b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2e-07, - "output_cost_per_token": 2.4e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-30b-a3b-instruct": { - "input_cost_per_token": 1.3e-07, - "output_cost_per_token": 5.2e-07, + "input_cost_per_token": 1.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 6e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67673,21 +68073,22 @@ "supports_web_search": true }, "openrouter/qwen/qwen3-vl-235b-a22b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 4e-07, - "output_cost_per_token": 4e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67712,32 +68113,33 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-max": { - "input_cost_per_token": 7.8e-07, - "input_cost_per_token_above_32k_tokens": 1.56e-06, - "cache_creation_input_token_cost_above_32k_tokens": 1.95e-06, - "cache_read_input_token_cost_above_32k_tokens": 3.12e-07, - "output_cost_per_token_above_32k_tokens": 7.8e-06, - "output_cost_per_token": 3.9e-06, - "cache_read_input_token_cost": 1.56e-07, "cache_creation_input_token_cost": 9.75e-07, - "input_cost_per_token_above_128k_tokens": 1.95e-06, - "output_cost_per_token_above_128k_tokens": 9.75e-06, - "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, "cache_creation_input_token_cost_above_128k_tokens": 2.4375e-06, + "cache_creation_input_token_cost_above_32k_tokens": 1.95e-06, + "cache_read_input_token_cost": 1.56e-07, + "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, + "cache_read_input_token_cost_above_32k_tokens": 3.12e-07, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 7.8e-07, + "input_cost_per_token_above_128k_tokens": 1.95e-06, + "input_cost_per_token_above_32k_tokens": 1.56e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 3.9e-06, + "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 7.8e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, - "supports_response_schema": true, - "supports_vision": false, "supports_prompt_caching": true, "supports_reasoning": false, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/deepseek/deepseek-v3.1-terminus": { @@ -67832,23 +68234,24 @@ "openrouter/qwen/qwen-plus-2025-07-28": { "cache_creation_input_token_cost": 3.25e-07, "cache_read_input_token_cost": 5.2e-08, + "deprecation_date": "2026-10-09", "input_cost_per_token": 2.6e-07, - "output_cost_per_token": 7.8e-07, "input_cost_per_token_above_256k_tokens": 7.8e-07, - "output_cost_per_token_above_256k_tokens": 2.34e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 7.8e-07, + "output_cost_per_token_above_256k_tokens": 2.34e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -67872,21 +68275,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-30b-a3b-thinking-2507": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2e-07, - "output_cost_per_token": 2.4e-06, "litellm_provider": "openrouter", "max_input_tokens": 81920, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -68195,21 +68599,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-8b": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.17e-07, - "output_cost_per_token": 4.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", + "output_cost_per_token": 4.55e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -68252,21 +68657,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-235b-a22b": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 4.55e-07, - "output_cost_per_token": 1.82e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", + "output_cost_per_token": 1.82e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -71335,6 +71741,23 @@ "output_cost_per_token": 0.0, "source": "https://openrouter.ai/typesafe/jev-1.13" }, + "openrouter/typesafe/jev-router": { + "input_cost_per_token": 0, + "output_cost_per_token": 0, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "source": "https://openrouter.ai/typesafe/jev-router", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_audio_input": true, + "supports_video_input": true + }, "typesafe/jev-1.13.0": { "input_cost_per_token": 4.2e-08, "litellm_provider": "typesafe", @@ -73019,14 +73442,14 @@ "supports_web_search": false }, "openrouter/inclusionai/ling-3.0-flash-vl": { - "cache_read_input_token_cost": 1.2e-08, - "input_cost_per_token": 6e-08, + "cache_read_input_token_cost": 4.2e-09, + "input_cost_per_token": 2.1e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 1.8e-07, + "output_cost_per_token": 6.16e-08, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -76537,6 +76960,23 @@ "supports_vision": true, "supports_web_search": false }, + "openrouter/perceptron/perceptron-mk1.5": { + "input_cost_per_token": 1.5e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 36864, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true + }, "vertex_ai/gemini-2.0-flash": { "deprecation_date": "2026-06-01", "input_cost_per_audio_token": 1e-06, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 6baa695433c..8befc99cad4 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -6855,12 +6855,13 @@ class MCPServerManager: if not tool_permissions: return {} expanded: Final = tuple( - (server_id, tuple(tools or ())) - for key, tools in tool_permissions.items() - for server_id in self.expand_permission_list([key]) + chain.from_iterable( + ((server_id, tuple(tools or ())) for server_id in self.expand_permission_list([key])) + for key, tools in tool_permissions.items() + ) ) return { - server_id: list(dict.fromkeys(tool for _, tools in group for tool in tools)) + server_id: list(dict.fromkeys(chain.from_iterable(tools for _, tools in group))) for server_id, group in groupby(sorted(expanded, key=itemgetter(0)), key=itemgetter(0)) } diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index f19a8055ae6..12d420141f1 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -18,7 +18,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias from fastapi import HTTPException, Request, status -from pydantic import BaseModel, TypeAdapter +from pydantic import BaseModel, TypeAdapter, ValidationError from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, Unpack import litellm @@ -5682,6 +5682,64 @@ async def _virtual_key_max_budget_alert_check( ) +TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY: Final = "team_member_max_budget_alert_emails" +_TEAM_MEMBER_ALERT_CONFIG_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object]) + + +def _is_valid_alert_threshold_pct(pct: str) -> bool: + return pct.isdigit() and len(pct) <= 3 and 1 <= int(pct) <= 100 + + +def _alert_recipients(raw: object) -> Sequence[str] | None: + if isinstance(raw, (str, Sequence)): + return _parse_email_list(raw) + return None + + +def _valid_alert_threshold_config(raw_config: object) -> Mapping[str, str | Sequence[object] | None] | None: + try: + config: Final = _TEAM_MEMBER_ALERT_CONFIG_ADAPTER.validate_python(raw_config) + except ValidationError: + return None + return MappingProxyType( + {pct: _alert_recipients(emails) for pct, emails in config.items() if _is_valid_alert_threshold_pct(pct)} + ) + + +def _team_member_max_budget_alert_check( + team_id: str, + team_alias: str | None, + team_metadata: Mapping[str, object] | None, + organization_id: str | None, + user_id: str, + user_email: str | None, + proxy_logging_obj: ProxyLogging, + spend: float, + max_budget: float, +) -> None: + raw_config: Final = team_metadata.get(TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY) if team_metadata else None + alert_email_config: Final = _merge_budget_alert_email_configs( + global_cfg=None, per_key_cfg=_valid_alert_threshold_config(raw_config) + ) + if not alert_email_config or spend <= 0: + return + min_pct: Final = min(int(pct) for pct in alert_email_config) + if spend < max_budget * (min_pct / 100.0): + return + call_info: Final = CallInfo( + spend=spend, + max_budget=max_budget, + user_id=user_id, + team_id=team_id, + team_alias=team_alias, + organization_id=organization_id, + user_email=user_email, + event_group=Litellm_EntityType.TEAM_MEMBER, + max_budget_alert_emails=alert_email_config, + ) + asyncio.create_task(proxy_logging_obj.budget_alerts(type="max_budget_alert", user_info=call_info)) + + async def _check_team_member_budget( team_object: LiteLLM_TeamTable | None, user_object: LiteLLM_UserTable | None, @@ -5747,7 +5805,22 @@ async def _check_team_member_budget( max_budget=team_member_budget, ) - if math.isfinite(team_member_budget) and team_member_spend >= team_member_budget: + if not math.isfinite(team_member_budget): + return + + _team_member_max_budget_alert_check( + team_id=team_object.team_id, + team_alias=team_object.team_alias, + team_metadata=team_object.metadata, + organization_id=team_object.organization_id, + user_id=valid_token.user_id, + user_email=user_object.user_email if user_object is not None else None, + proxy_logging_obj=proxy_logging_obj, + spend=team_member_spend, + max_budget=team_member_budget, + ) + + if team_member_spend >= team_member_budget: raise litellm.BudgetExceededError( current_cost=team_member_spend, max_budget=team_member_budget, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 22c3a248b9d..e3ce9bcd850 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -50,6 +50,7 @@ from litellm.proxy.auth.auth_checks import ( _get_user_role, _is_model_cost_zero, _is_user_proxy_admin, + _team_member_max_budget_alert_check, _virtual_key_max_budget_alert_check, _virtual_key_max_budget_check, _virtual_key_soft_budget_check, @@ -2287,6 +2288,19 @@ async def _user_api_key_auth_builder( max_budget=team_member_budget, ) if team_member_spend >= team_member_budget: + # common_checks sends this alert on requests that get past here, so only the + # request rejected here sends it from the builder. + _team_member_max_budget_alert_check( + team_id=_team_id, + team_alias=valid_token.team_alias, + team_metadata=valid_token.team_metadata, + organization_id=valid_token.org_id, + user_id=_user_id, + user_email=user_obj.user_email if user_obj is not None else None, + proxy_logging_obj=proxy_logging_obj, + spend=team_member_spend, + max_budget=team_member_budget, + ) _entity_id: Final = f"{valid_token.user_id}:{valid_token.team_id}" raise litellm.BudgetExceededError( current_cost=team_member_spend, diff --git a/litellm/proxy/common_utils/model_listing_utils.py b/litellm/proxy/common_utils/model_listing_utils.py index 8958fb20918..0c702ac4139 100644 --- a/litellm/proxy/common_utils/model_listing_utils.py +++ b/litellm/proxy/common_utils/model_listing_utils.py @@ -14,6 +14,7 @@ import re from collections.abc import Container, Mapping, Sequence from dataclasses import dataclass from functools import reduce +from itertools import chain from types import MappingProxyType from typing import TYPE_CHECKING, Final, cast @@ -191,7 +192,7 @@ def alias_map(aliases: object) -> Mapping[str, str]: def _alias_names(alias_maps: Sequence[Mapping[str, str]]) -> tuple[str, ...]: - return tuple(dict.fromkeys(alias for aliases in alias_maps for alias in aliases)) + return tuple(dict.fromkeys(chain.from_iterable(alias_maps))) def _rewrite(model_id: str, alias_maps: Sequence[Mapping[str, str]]) -> str | None: diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index e097debde77..0178465739b 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -474,7 +474,9 @@ class _ProxyDBLogger(CustomLogger): spend_log_error("Error in tracking cost callback - %s", str(e), exc=e) @staticmethod - async def _enrich_failure_metadata_unless_db_stalled(metadata: dict, original_exception: Exception) -> dict: + async def _enrich_failure_metadata_unless_db_stalled( + metadata: dict[str, object], original_exception: Exception + ) -> dict[str, object]: if isinstance(original_exception, DBLookupDeadlineExceeded): return metadata return await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata) diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 08e8fff8f1d..965cded59c4 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -215,17 +215,23 @@ def _meta_with_user_details( return updated +def _user_id_needing_details(api_key: str, meta: KeyMetadataDict) -> str | None: + user_id: Final = meta.get("user_id") + if not isinstance(user_id, str) or not user_id: + return None + if meta.get("user_email") and not (_is_cli_session_key(api_key) and not meta.get("team_id")): + return None + return user_id + + async def attach_user_details( prisma_client: PrismaClient, recovered: Mapping[str, KeyMetadataDict], ) -> Mapping[str, KeyMetadataDict]: needing_details: Final = frozenset( user_id - for api_key, meta in recovered.items() - for user_id in (meta.get("user_id"),) - if isinstance(user_id, str) - and user_id - and (not meta.get("user_email") or (_is_cli_session_key(api_key) and not meta.get("team_id"))) + for user_id in (_user_id_needing_details(api_key, meta) for api_key, meta in recovered.items()) + if user_id is not None ) details: Final = await _details_for_user_ids(prisma_client, needing_details) if not details: diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 1ef39775bd3..9f537d24eaa 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -265,7 +265,7 @@ def _mid_stream_fallback_eligible(mapped_exception: Exception) -> bool: return not isinstance(status_code, int) or status_code >= 500 or status_code == 429 -_PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: Final = frozenset({"response.created", "response.in_progress", "response.queued"}) +PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: Final = frozenset({"response.created", "response.in_progress", "response.queued"}) class BaseResponsesAPIStreamingIterator: @@ -885,7 +885,7 @@ class BaseResponsesAPIStreamingIterator: def _note_yielded_event(self, event: ResponsesAPIStreamingResponse) -> None: self._yielded_first_chunk = True - if event.type not in _PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: + if event.type not in PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: self._output_started = True def _fallback_error(self, original: Exception) -> MidStreamFallbackError: diff --git a/litellm/router.py b/litellm/router.py index 023b99cd64e..6f416c416c0 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -614,6 +614,17 @@ def _anthropic_stream_commits_now(chunk: object, has_generated_content: bool, bu return is_anthropic_content_delta_chunk(chunk) or buffered_chunk_count >= MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS +MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS: Final = 200 + + +def _responses_stream_holds_event(item: object, held_event_count: int) -> bool: + from litellm.responses.streaming_iterator import PRE_OUTPUT_LIFECYCLE_EVENT_TYPES + + if held_event_count >= MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS: + return False + return getattr(item, "type", None) in PRE_OUTPUT_LIFECYCLE_EVENT_TYPES + + class FallbackAwareAnthropicMessagesStream: """ Bare async generators can't carry the `_hidden_params` attribute the @@ -3332,100 +3343,140 @@ class Router: await self._async_generator.aclose() async def stream_with_fallbacks(): - fallback_response = None + held_lifecycle_events: tuple[object, ...] = () # rebind-ok: flushed at first output, dropped on fallback try: async for item in source_iterator: + if _responses_stream_holds_event(item, len(held_lifecycle_events)): + held_lifecycle_events = (*held_lifecycle_events, item) + continue + for held_event in held_lifecycle_events: + yield held_event + held_lifecycle_events = () yield item + for held_event in held_lifecycle_events: + yield held_event except MidStreamFallbackError as e: - partial_usage: Final = Router._extract_partial_responses_usage(source_iterator) - try: - model_group: Final = cast(str, initial_kwargs.get("model")) - fallbacks: Final[list | None] = initial_kwargs.get("fallbacks", self.fallbacks) - context_window_fallbacks: Final[list | None] = initial_kwargs.get( - "context_window_fallbacks", self.context_window_fallbacks + async with contextlib.aclosing( + self._aresponses_fallback_attempt( + e, source_iterator, initial_kwargs, wrapper.adopt_fallback_headers, held_lifecycle_events ) - content_policy_fallbacks: Final[list | None] = initial_kwargs.get( - "content_policy_fallbacks", self.content_policy_fallbacks - ) - initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_responses_attempt - if e.is_pre_first_chunk or not e.generated_content: - # No content generated before the error — retry with the - # original input. Adding a continuation prompt would - # waste tokens and confuse the model. - pass - else: - initial_kwargs["input"] = Router._build_responses_continuation_input( - initial_kwargs.get("input"), - e.generated_content, - ) - # The Responses-API path stores observability metadata - # under "litellm_metadata" (not the default "metadata") — - # see _ageneric_api_call_with_fallbacks. Mirroring that - # here ensures model_group, model_group_alias, and trace - # ids land in the same key litellm.aresponses reads from. - self._update_kwargs_before_fallbacks( - model=model_group, - kwargs=initial_kwargs, - metadata_variable_name="litellm_metadata", - ) - # The content-policy dispatch branch matches on the trigger's own type, so a refusal's - # MidStreamFallbackError envelope is unwrapped here or the wrong fallback list is consulted. - fallback_trigger: Final[Exception] = ( - e.original_exception - if isinstance(e.original_exception, litellm.ContentPolicyViolationError) - else e - ) - fallback_response = await self.async_function_with_fallbacks_common_utils( - e=fallback_trigger, - disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), - fallbacks=fallbacks, - context_window_fallbacks=context_window_fallbacks, - content_policy_fallbacks=content_policy_fallbacks, - model_group=model_group, - args=(), - kwargs=initial_kwargs, - include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True, - ) - - prepared_fallback_hidden_params = wrapper.adopt_fallback_headers(fallback_response) - if hasattr(fallback_response, "__aiter__"): - async for fallback_item in fallback_response: - Router._apply_fallback_hidden_params_to_item(fallback_item, prepared_fallback_hidden_params) - if partial_usage is not None: - Router._combine_responses_fallback_usage(fallback_item, partial_usage) - yield fallback_item - else: - yield fallback_response - except Exception as fallback_error: - verbose_router_logger.error("Responses streaming fallback also failed: %s", fallback_error) - if ( - isinstance(fallback_error, MidStreamFallbackError) - and fallback_error.original_exception is not None - ): - raise fallback_error.original_exception from fallback_error - raise fallback_error + ) as fallback_stream: + async for fallback_item in fallback_stream: + yield fallback_item + except Exception: + for held_event in held_lifecycle_events: + yield held_event + raise finally: with anyio.CancelScope(shield=True): if hasattr(source_iterator, "aclose"): try: await source_iterator.aclose() - except BaseException as exc: + except Exception as exc: verbose_router_logger.debug( "stream_with_fallbacks(aresponses): error closing source: %s", exc, ) - if fallback_response is not None and hasattr(fallback_response, "aclose"): - try: - await fallback_response.aclose() - except BaseException as exc: - verbose_router_logger.debug( - "stream_with_fallbacks(aresponses): error closing fallback: %s", - exc, - ) wrapper: Final = FallbackResponsesStreamWrapper(stream_with_fallbacks()) return wrapper + async def _aresponses_fallback_attempt( + self, + e: "MidStreamFallbackError", + source_iterator: "BaseResponsesAPIStreamingIterator", + initial_kwargs: dict[str, Any], # mutable-ok: mutated in-place before re-entering the fallback chain + adopt_headers: Callable[[object], tuple[dict[str, object], dict[str, object]]], # mutable-ok: hidden params + held_lifecycle_events: tuple[object, ...], + ) -> AsyncGenerator[object, None]: + """ + Re-enters the Router's fallback chain for a mid-stream Responses API error and yields + whatever the fallback attempt produces. The lifecycle events the primary stream held + back reach the client only when no fallback lands, so the client sees exactly one + response announced, the one whose id completes. Split out of + _aresponses_streaming_iterator to keep each function's cyclomatic complexity within + the repo's C901 budget. + """ + from litellm.exceptions import MidStreamFallbackError + + partial_usage: Final = Router._extract_partial_responses_usage(source_iterator) + fallback_response = None # rebind-ok: pre-init so finally can close it if a fallback was actually attempted + fallback_yielded = False # rebind-ok: flipped on the first fallback item so a fallback that dies before its first event still replays the primary's held announcement + try: + model_group: Final = cast(str, initial_kwargs.get("model")) # cast-ok: model group + fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the common_utils list|None param + "fallbacks", self.fallbacks + ) + context_window_fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the param below + "context_window_fallbacks", self.context_window_fallbacks + ) + content_policy_fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the param below + "content_policy_fallbacks", self.content_policy_fallbacks + ) + initial_kwargs["original_function"] = ( # rebind-ok: the fallback chain re-enters on the same kwargs + self._ageneric_api_call_with_fallbacks_responses_attempt + ) + if e.generated_content and not e.is_pre_first_chunk: + initial_kwargs["input"] = Router._build_responses_continuation_input( # rebind-ok: fallback hop input + initial_kwargs.get("input"), + e.generated_content, + ) + # The Responses-API path stores observability metadata + # under "litellm_metadata" (not the default "metadata") — + # see _ageneric_api_call_with_fallbacks. Mirroring that + # here ensures model_group, model_group_alias, and trace + # ids land in the same key litellm.aresponses reads from. + self._update_kwargs_before_fallbacks( + model=model_group, + kwargs=initial_kwargs, + metadata_variable_name="litellm_metadata", + ) + # The content-policy dispatch branch matches on the trigger's own type, so a refusal's + # MidStreamFallbackError envelope is unwrapped here or the wrong fallback list is consulted. + fallback_trigger: Final[Exception] = ( + e.original_exception if isinstance(e.original_exception, litellm.ContentPolicyViolationError) else e + ) + fallback_response = await self.async_function_with_fallbacks_common_utils( # rebind-ok: set on success + e=fallback_trigger, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), + fallbacks=fallbacks, + context_window_fallbacks=context_window_fallbacks, + content_policy_fallbacks=content_policy_fallbacks, + model_group=model_group, + args=(), + kwargs=initial_kwargs, + include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True, + ) + prepared_fallback_hidden_params: Final = adopt_headers(fallback_response) + if hasattr(fallback_response, "__aiter__"): + async for fallback_item in fallback_response: + Router._apply_fallback_hidden_params_to_item(fallback_item, prepared_fallback_hidden_params) + if partial_usage is not None: + Router._combine_responses_fallback_usage(fallback_item, partial_usage) + fallback_yielded = True + yield fallback_item + else: + fallback_yielded = True # rebind-ok: see the pre-init above + yield fallback_response + except Exception as fallback_error: + verbose_router_logger.error("Responses streaming fallback also failed: %s", fallback_error) + if not fallback_yielded: + for held_event in held_lifecycle_events: + yield held_event + if isinstance(fallback_error, MidStreamFallbackError) and fallback_error.original_exception is not None: + raise fallback_error.original_exception from fallback_error + raise + finally: + if fallback_response is not None and hasattr(fallback_response, "aclose"): + with anyio.CancelScope(shield=True): + try: + await fallback_response.aclose() + except Exception as exc: + verbose_router_logger.debug( + "stream_with_fallbacks(aresponses): error closing fallback: %s", + exc, + ) + def _completion_streaming_iterator( self, model_response: CustomStreamWrapper, diff --git a/litellm/rust_bridge/callbacks_legacy_python.py b/litellm/rust_bridge/callbacks_legacy_python.py index 6bbf2ffed6b..e39324d3348 100644 --- a/litellm/rust_bridge/callbacks_legacy_python.py +++ b/litellm/rust_bridge/callbacks_legacy_python.py @@ -22,7 +22,6 @@ from typing import ( if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging - from litellm.types.utils import CredentialItem class MetadataUpdater(Protocol): @@ -72,21 +71,6 @@ def _claim_budget_reservation(call_setup: CallSetup, asynchronous: bool) -> Call return call_setup -def check_limits(kwargs: Mapping[str, object]) -> None: - from litellm import ( - BudgetExceededError, - _current_cost, # pyright: ignore[reportPrivateUsage] # shared SDK budget counter has no public accessor - max_budget, - num_retries_per_request, - ) - from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit - - if max_budget and _current_cost > max_budget: - raise BudgetExceededError(current_cost=_current_cost, max_budget=max_budget) - if max_retries_per_request_hit(kwargs, num_retries_per_request): - raise RuntimeError("Max retries per request hit!") - - def finalize( response: object, logger: Logging, @@ -299,22 +283,6 @@ def is_internal_call() -> bool: return internal.get() -def credential_list() -> list[CredentialItem]: - from litellm import credential_list as credentials - - return credentials - - -def warn_unknown_credential(name: str, loaded: int) -> None: - from litellm._logging import verbose_logger - - verbose_logger.warning( - "litellm_credential_name=%s matched none of the %d loaded credentials; the request runs without it", - name, - loaded, - ) - - def before_deployment_call(kwargs: dict[str, object], call_type: str) -> Awaitable[object]: from litellm import utils diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index 6e455817194..46acb94c958 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -109,7 +109,7 @@ RULES: Final[Rules] = ( RouteRule(Route.CHAT_COMPLETIONS, Rollout.PYTHON_ONLY), RouteRule(Route.EMBEDDINGS, Rollout.PYTHON_ONLY), RouteRule(Route.OCR, Rollout.RUST_REQUIRED), - RouteRule(Route.MESSAGES, Rollout.PYTHON_ONLY, providers=frozenset({"anthropic"})), + RouteRule(Route.MESSAGES, Rollout.RUST_OPT_IN, providers=frozenset({"anthropic"})), RouteRule(Route.MESSAGES, Rollout.PYTHON_ONLY), RouteRule(Route.RESPONSES, Rollout.PYTHON_ONLY), RouteRule(Route.TOKEN_COUNTER, Rollout.PYTHON_ONLY), diff --git a/litellm/rust_bridge/preflight.py b/litellm/rust_bridge/preflight.py new file mode 100644 index 00000000000..e030382bfdc --- /dev/null +++ b/litellm/rust_bridge/preflight.py @@ -0,0 +1,45 @@ +"""The SDK request policy the native driver runs before a route's host projects. + +These are the `@client` prologue steps after `function_setup` and the deployment hook: +credential-name inheritance and the budget and retry-count limits. Rust owns the +inheritance itself; it borrows only the globals below. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from litellm.types.utils import CredentialItem + + +def credential_list() -> list[CredentialItem]: + from litellm import credential_list as credentials + + return credentials + + +def warn_unknown_credential(name: str, loaded: int) -> None: + from litellm._logging import verbose_logger + + verbose_logger.warning( + "litellm_credential_name=%s matched none of the %d loaded credentials; the request runs without it", + name, + loaded, + ) + + +def check_limits(kwargs: Mapping[str, object]) -> None: + from litellm import ( + BudgetExceededError, + _current_cost, # pyright: ignore[reportPrivateUsage] # shared SDK budget counter has no public accessor + max_budget, + num_retries_per_request, + ) + from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit + + if max_budget and _current_cost > max_budget: + raise BudgetExceededError(current_cost=_current_cost, max_budget=max_budget) + if max_retries_per_request_hit(kwargs, num_retries_per_request): + raise RuntimeError("Max retries per request hit!") diff --git a/litellm/secret_managers/cyberark_secret_manager.py b/litellm/secret_managers/cyberark_secret_manager.py index b28e15c4446..f8e17488167 100644 --- a/litellm/secret_managers/cyberark_secret_manager.py +++ b/litellm/secret_managers/cyberark_secret_manager.py @@ -1,3 +1,4 @@ +import asyncio import base64 import os from typing import Any, Final @@ -10,6 +11,7 @@ import litellm from litellm._logging import verbose_logger from litellm.caching import InMemoryCache from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, _get_httpx_client, get_async_httpx_client, httpxSpecialProvider, @@ -20,6 +22,9 @@ from litellm.rust_bridge.secret_manager import resolve_native_provider_reader, r from .base_secret_manager import BaseSecretManager, raise_if_unsafe_secret_name from .main import str_to_bool +CYBERARK_POLICY_LOAD_ATTEMPTS: Final = 5 +CYBERARK_POLICY_LOAD_RETRY_DELAY_SECONDS: Final = 0.2 + class CyberArkSecretManager(BaseSecretManager): def __init__(self): @@ -30,6 +35,7 @@ class CyberArkSecretManager(BaseSecretManager): self.conjur_account = os.getenv("CYBERARK_ACCOUNT", "default") self.conjur_username = os.getenv("CYBERARK_USERNAME", "admin") self.conjur_api_key = os.getenv("CYBERARK_API_KEY", "") + self._policy_load_lock: Final = asyncio.Lock() # Optional config for certificate-based auth self.tls_cert_path = os.getenv("CYBERARK_CLIENT_CERT", "") @@ -118,7 +124,7 @@ class CyberArkSecretManager(BaseSecretManager): token: Final = self._authenticate() return {"Authorization": f'Token token="{token}"'} - def _ensure_variable_exists(self, secret_name: str) -> None: + async def _ensure_variable_exists(self, secret_name: str, async_client: AsyncHTTPHandler) -> None: """ Ensure a variable exists in CyberArk Conjur by creating a policy entry if needed. @@ -134,27 +140,33 @@ class CyberArkSecretManager(BaseSecretManager): policy_yaml: Final = f"- !variable {quoted_name}\n" try: - client: Final = _get_httpx_client(params={"ssl_verify": self.ssl_verify}) - resp: Final = client.client.post( - policy_url, - headers={ - **self._get_request_headers(), - "Content-Type": "application/x-yaml", - }, - content=policy_yaml, - ) - resp.raise_for_status() - verbose_logger.debug("Created policy entry for variable: %s", secret_name) - except httpx.HTTPStatusError as e: - # Variable might already exist, which is fine - if e.response.status_code in [409, 422]: - verbose_logger.debug("Variable %s already exists or policy conflict (expected)", secret_name) - else: - verbose_logger.warning( - "Could not ensure variable exists: %s - %s", e.response.status_code, e.response.text - ) + async with self._policy_load_lock: + resp: Final = await self._load_variable_policy(async_client, policy_url, policy_yaml) except Exception as e: verbose_logger.warning("Error ensuring variable exists: %s", e) + return + if resp.is_success: + verbose_logger.debug("Created policy entry for variable: %s", secret_name) + elif resp.status_code == 422: + verbose_logger.debug("Variable %s policy was rejected as unprocessable", secret_name) + else: + verbose_logger.warning("Could not ensure variable exists: %s - %s", resp.status_code, resp.text) + + async def _load_variable_policy( + self, async_client: AsyncHTTPHandler, policy_url: str, policy_yaml: str, attempt: int = 0 + ) -> httpx.Response: + resp: Final = await async_client.client.post( + policy_url, + headers={ + **self._get_request_headers(), + "Content-Type": "application/x-yaml", + }, + content=policy_yaml, + ) + if resp.status_code != 409 or attempt + 1 == CYBERARK_POLICY_LOAD_ATTEMPTS: + return resp + await asyncio.sleep(CYBERARK_POLICY_LOAD_RETRY_DELAY_SECONDS * (1 << attempt)) + return await self._load_variable_policy(async_client, policy_url, policy_yaml, attempt + 1) def get_url(self, secret_name: str) -> str: """ @@ -303,7 +315,7 @@ class CyberArkSecretManager(BaseSecretManager): try: # Ensure the variable exists in the policy first - self._ensure_variable_exists(secret_name) + await self._ensure_variable_exists(secret_name, async_client) # Now set the secret value url: Final = self.get_url(secret_name) diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index 83a42c235f9..f5ba9ebd3da 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -3,6 +3,7 @@ models and KWARG_ARTIFACTS into all_litellm_params.""" from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequence from dataclasses import dataclass, field, fields, is_dataclass +from itertools import chain from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal, TypeAlias @@ -359,6 +360,6 @@ def owned_wire_names(root: type) -> tuple[str, ...]: return tuple(names()) -OWNED_KWARG_NAMES: Final = tuple(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root)) +OWNED_KWARG_NAMES: Final = tuple(chain.from_iterable(owned_wire_names(root) for root in LITELLM_OWNED_ROOTS)) AGENTIC_LOOP_KWARG_NAMES: Final = (*wire_names(AgenticLoopState), *wire_names(AgenticLoopOptions)) BEDROCK_BATCH_KWARG_NAMES: Final = wire_names(BedrockBatchConnection) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index cd336c9b989..2e518af4da4 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -302,6 +302,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): cache_read_input_token_cost_above_200k_tokens_batches: ReadOnly[float | None] cache_read_input_token_cost_above_272k_tokens_batches: ReadOnly[float | None] cache_creation_input_token_cost_batches: ReadOnly[float | None] + cache_creation_input_token_cost_above_200k_tokens_batches: ReadOnly[float | None] cache_creation_input_token_cost_above_272k_tokens_batches: ReadOnly[float | None] # Smallest prefix this model will actually cache, whatever caching mechanism its provider uses. # Absent means the provider-agnostic default applies; see MINIMUM_PROMPT_CACHE_TOKEN_COUNT. @@ -3735,6 +3736,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams): cache_read_input_token_cost_above_200k_tokens_batches: float | None = None cache_read_input_token_cost_above_272k_tokens_batches: float | None = None cache_creation_input_token_cost_batches: float | None = None + cache_creation_input_token_cost_above_200k_tokens_batches: float | None = None cache_creation_input_token_cost_above_272k_tokens_batches: float | None = None cache_read_input_audio_token_cost: float | None = None cache_read_input_image_token_cost: float | None = None diff --git a/litellm/utils.py b/litellm/utils.py index be4388802f9..42da2e2a7b7 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6167,6 +6167,9 @@ def _get_model_info_helper( "cache_read_input_token_cost_above_272k_tokens_batches" ), cache_creation_input_token_cost_batches=_model_info.get("cache_creation_input_token_cost_batches"), + cache_creation_input_token_cost_above_200k_tokens_batches=_model_info.get( + "cache_creation_input_token_cost_above_200k_tokens_batches" + ), cache_creation_input_token_cost_above_272k_tokens_batches=_model_info.get( "cache_creation_input_token_cost_above_272k_tokens_batches" ), diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 792e7316c67..4b88072b479 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -1279,8 +1279,11 @@ "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "anthropic.claude-mythos-preview": { - "input_cost_per_token": 0, - "output_cost_per_token": 0, + "cache_creation_input_token_cost": 3.4375e-05, + "cache_creation_input_token_cost_above_1hr": 5.5e-05, + "cache_read_input_token_cost": 2.75e-06, + "input_cost_per_token": 2.75e-05, + "output_cost_per_token": 0.0001375, "litellm_provider": "bedrock", "max_input_tokens": 1000000, "max_output_tokens": 128000, @@ -1289,10 +1292,11 @@ "thinking_always_on": true, "supports_function_calling": true, "supports_vision": true, - "supports_prompt_caching": false, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_tool_choice": true, - "supports_output_config": true + "supports_output_config": true, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "global.anthropic.claude-opus-4-7": { "bedrock_converse_supports_strict_tools": false, @@ -7773,27 +7777,27 @@ "output_cost_per_token_above_272k_tokens_batches": 0.000135 }, "azure/gpt-5.6": { - "cache_creation_input_token_cost": 6.25e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.25e-05, - "cache_creation_input_token_cost_priority": 1.25e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.5e-05, - "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1e-06, - "cache_read_input_token_cost_priority": 1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2e-06, - "input_cost_per_token": 5e-06, - "input_cost_per_token_above_272k_tokens": 1e-05, - "input_cost_per_token_priority": 1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2e-05, + "cache_creation_input_token_cost": 5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1e-05, + "cache_creation_input_token_cost_priority": 1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2e-05, + "cache_read_input_token_cost": 4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8e-07, + "cache_read_input_token_cost_priority": 8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.6e-06, + "input_cost_per_token": 4e-06, + "input_cost_per_token_above_272k_tokens": 8e-06, + "input_cost_per_token_priority": 8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.6e-05, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3e-05, - "output_cost_per_token_above_272k_tokens": 4.5e-05, - "output_cost_per_token_priority": 6e-05, - "output_cost_per_token_above_272k_tokens_priority": 9e-05, + "output_cost_per_token": 2e-05, + "output_cost_per_token_above_272k_tokens": 3e-05, + "output_cost_per_token_priority": 4e-05, + "output_cost_per_token_above_272k_tokens_priority": 6e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -8579,27 +8583,27 @@ "supports_web_search": true }, "azure/us/gpt-5.6": { - "cache_creation_input_token_cost": 6.875e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.375e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.75e-05, - "cache_creation_input_token_cost_priority": 1.375e-05, - "cache_read_input_token_cost": 5.5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1.1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2.2e-06, - "cache_read_input_token_cost_priority": 1.1e-06, - "input_cost_per_token": 5.5e-06, - "input_cost_per_token_above_272k_tokens": 1.1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2.2e-05, - "input_cost_per_token_priority": 1.1e-05, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1.1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2.2e-05, + "cache_creation_input_token_cost_priority": 1.1e-05, + "cache_read_input_token_cost": 4.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8.8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.76e-06, + "cache_read_input_token_cost_priority": 8.8e-07, + "input_cost_per_token": 4.4e-06, + "input_cost_per_token_above_272k_tokens": 8.8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.76e-05, + "input_cost_per_token_priority": 8.8e-06, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3.3e-05, - "output_cost_per_token_above_272k_tokens": 4.95e-05, - "output_cost_per_token_above_272k_tokens_priority": 9.9e-05, - "output_cost_per_token_priority": 6.6e-05, + "output_cost_per_token": 2.2e-05, + "output_cost_per_token_above_272k_tokens": 3.3e-05, + "output_cost_per_token_above_272k_tokens_priority": 6.6e-05, + "output_cost_per_token_priority": 4.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -8985,27 +8989,27 @@ "supports_web_search": true }, "azure/eu/gpt-5.6": { - "cache_creation_input_token_cost": 6.875e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.375e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.75e-05, - "cache_creation_input_token_cost_priority": 1.375e-05, - "cache_read_input_token_cost": 5.5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1.1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2.2e-06, - "cache_read_input_token_cost_priority": 1.1e-06, - "input_cost_per_token": 5.5e-06, - "input_cost_per_token_above_272k_tokens": 1.1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2.2e-05, - "input_cost_per_token_priority": 1.1e-05, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1.1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2.2e-05, + "cache_creation_input_token_cost_priority": 1.1e-05, + "cache_read_input_token_cost": 4.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8.8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.76e-06, + "cache_read_input_token_cost_priority": 8.8e-07, + "input_cost_per_token": 4.4e-06, + "input_cost_per_token_above_272k_tokens": 8.8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.76e-05, + "input_cost_per_token_priority": 8.8e-06, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3.3e-05, - "output_cost_per_token_above_272k_tokens": 4.95e-05, - "output_cost_per_token_above_272k_tokens_priority": 9.9e-05, - "output_cost_per_token_priority": 6.6e-05, + "output_cost_per_token": 2.2e-05, + "output_cost_per_token_above_272k_tokens": 3.3e-05, + "output_cost_per_token_above_272k_tokens_priority": 6.6e-05, + "output_cost_per_token_priority": 4.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -11709,8 +11713,8 @@ "input_cost_per_token": 1.75e-06, "litellm_provider": "azure_ai", "mode": "image_generation", - "output_cost_per_image": 0.0338, - "output_cost_per_image_token": 3.3e-05, + "output_cost_per_image": 0.02, + "output_cost_per_image_token": 1.95e-05, "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", "supported_endpoints": [ "/v1/images/generations", @@ -14640,14 +14644,18 @@ "claude-haiku-4-5-20251001": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_creation_input_token_cost_batches": 6.25e-07, "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 5e-06, + "output_cost_per_token_batches": 2.5e-06, "supports_assistant_prefill": true, "supports_function_calling": true, "supports_native_structured_output": true, @@ -14663,14 +14671,18 @@ "claude-haiku-4-5": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_creation_input_token_cost_batches": 6.25e-07, "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 5e-06, + "output_cost_per_token_batches": 2.5e-06, "supports_assistant_prefill": true, "supports_function_calling": true, "supports_native_structured_output": true, @@ -14693,13 +14705,21 @@ "input_cost_per_token_above_200k_tokens": 6e-06, "output_cost_per_token_above_200k_tokens": 2.25e-05, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_creation_input_token_cost_above_200k_tokens_batches": 3.75e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost_above_200k_tokens": 6e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, + "input_cost_per_token_above_200k_tokens_batches": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_above_200k_tokens_batches": 1.125e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14727,13 +14747,21 @@ "input_cost_per_token_above_200k_tokens": 6e-06, "output_cost_per_token_above_200k_tokens": 2.25e-05, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_creation_input_token_cost_above_200k_tokens_batches": 3.75e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost_above_200k_tokens": 6e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, + "input_cost_per_token_above_200k_tokens_batches": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_above_200k_tokens_batches": 1.125e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14757,14 +14785,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, + "cache_creation_input_token_cost_batches": 1.25e-06, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14797,14 +14829,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 3.75e-06, "cache_creation_input_token_cost_above_1hr": 6e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, "input_cost_per_token": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14865,14 +14901,18 @@ "claude-opus-4-5-20251101": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14895,14 +14935,18 @@ "claude-opus-4-5": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14927,14 +14971,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14966,14 +15014,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15004,14 +15056,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15044,14 +15100,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15083,14 +15143,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 1e-06, + "cache_read_input_token_cost_batches": 5e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15123,14 +15187,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_batches": 1.25e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15164,14 +15232,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 5e-06, "cache_creation_input_token_cost_above_1hr": 8e-06, + "cache_creation_input_token_cost_batches": 2.5e-06, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 4e-06, + "input_cost_per_token_batches": 2e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2e-05, + "output_cost_per_token_batches": 1e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15207,14 +15279,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15250,14 +15326,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15756,6 +15836,16 @@ "supports_function_calling": true, "supports_tool_choice": true }, + "c4ai-aya-expanse-32b": { + "input_cost_per_token": 5e-07, + "output_cost_per_token": 1.5e-06, + "litellm_provider": "cohere_chat", + "max_input_tokens": 128000, + "max_output_tokens": 4000, + "max_tokens": 4000, + "mode": "chat", + "source": "https://docs.cohere.com/docs/models" + }, "command-a-plus-05-2026": { "input_cost_per_token": 0.0, "litellm_provider": "cohere_chat", @@ -19104,6 +19194,41 @@ "supports_tool_choice": true, "supports_vision": true }, + "databricks/databricks-claude-opus-5-5": { + "cache_creation_input_token_cost": 5.00003e-06, + "cache_creation_input_token_cost_above_1hr": 8.00002e-06, + "cache_read_input_token_cost": 1.9999e-07, + "input_cost_per_token": 4.00001e-06, + "input_dbu_cost_per_token": 5.7143e-05, + "litellm_provider": "databricks", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Costs per token are the published Global DBU rates times $0.070 per DBU. The '*_dbu_cost_per_token' fields are provided for reference; cost calculation reads the dollar '*_cost_per_token' fields." + }, + "mode": "chat", + "output_cost_per_token": 1.999998e-05, + "output_dbu_cost_per_token": 0.000285714, + "prompt_cache_min_tokens": 512, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_adaptive_thinking": true, + "supports_anthropic_thinking_payload": true, + "supports_assistant_prefill": false, + "supports_forced_tool_use": false, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_output_config": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "thinking_always_on": true + }, "databricks/databricks-claude-sonnet-4": { "cache_creation_input_token_cost": 3.74997e-06, "cache_read_input_token_cost": 3.0002e-07, @@ -22145,11 +22270,11 @@ } }, "bing_grounding/search": { - "input_cost_per_query": 0.035, + "input_cost_per_query": 0.014, "litellm_provider": "bing_grounding", "mode": "search", "metadata": { - "notes": "Grounding with Bing Search (G1 SKU): $35 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment." + "notes": "Grounding with Bing Search (G1 SKU): $14 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment." } }, "tinyfish/search": { @@ -28036,6 +28161,7 @@ "tpm": 10000000 }, "gemini/gemini-3-pro-image-preview": { + "deprecation_date": "2026-06-25", "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, @@ -28083,6 +28209,7 @@ "supports_reasoning": false }, "gemini/gemini-3.1-flash-image-preview": { + "deprecation_date": "2026-06-25", "input_cost_per_token": 5e-07, "input_cost_per_token_batches": 2.5e-07, "litellm_provider": "gemini", @@ -28130,6 +28257,7 @@ "cache_read_input_token_cost_batches": 1.25e-08, "cache_read_input_token_cost_flex": 1.25e-08, "cache_read_input_token_cost_priority": 4.5e-08, + "deprecation_date": "2026-05-25", "input_cost_per_audio_token": 5e-07, "input_cost_per_token": 2.5e-07, "input_cost_per_token_batches": 1.25e-07, @@ -30551,7 +30679,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.6-luna": { "litellm_provider": "chatgpt", @@ -30567,7 +30699,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.6-sol": { "litellm_provider": "chatgpt", @@ -30583,7 +30719,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.6-terra": { "litellm_provider": "chatgpt", @@ -30599,7 +30739,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.4": { "litellm_provider": "chatgpt", @@ -30614,7 +30758,12 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "default_reasoning_effort": "none", + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.4-pro": { "litellm_provider": "chatgpt", @@ -30628,7 +30777,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": false, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.3-codex": { "litellm_provider": "chatgpt", @@ -30642,7 +30795,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": false, + "supports_xhigh_reasoning_effort": false, + "supports_minimal_reasoning_effort": true }, "chatgpt/gpt-5.3-codex-spark": { "litellm_provider": "chatgpt", @@ -30686,7 +30843,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.2-codex": { "litellm_provider": "chatgpt", @@ -30700,7 +30858,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.2": { "litellm_provider": "chatgpt", @@ -30715,7 +30874,12 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "default_reasoning_effort": "none", + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.1-codex-max": { "litellm_provider": "chatgpt", @@ -30729,7 +30893,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.1-codex-mini": { "litellm_provider": "chatgpt", @@ -30743,7 +30908,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "gigachat/GigaChat-2": { "input_cost_per_token": 0.0, @@ -39214,6 +39380,17 @@ "supports_function_calling": true, "supports_reasoning": true }, + "nebius/deepseek-ai/DeepSeek-V4.1-Flash": { + "input_cost_per_token": 3e-07, + "litellm_provider": "nebius", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://tokenfactory.nebius.com/endpoints?modals=endpoint-details&model-id=deepseek-ai/DeepSeek-V4.1-Flash", + "supports_vision": true + }, "nebius/MiniMaxAI/MiniMax-M2.5": { "max_tokens": 196608, "max_input_tokens": 196608, @@ -41477,65 +41654,63 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 8.44944e-07, + "cache_read_input_token_cost": 3.828e-08, + "input_cost_per_token": 4.5936e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.689888e-06, + "output_cost_per_token": 9.1872e-07, "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, "supports_function_calling": true, + "supports_pdf_input": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 7.0412e-08, - "supports_audio_input": false, - "supports_pdf_input": false, "supports_vision": false, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "input_cost_per_token": 3e-07, - "output_cost_per_token": 1.2e-06, - "cache_read_input_token_cost": 6e-09, + "cache_read_input_token_cost": 4.2e-09, + "input_cost_per_token": 1.4e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"00:00-01:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"04:00-06:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"10:00-00:00"}],"input_cost_per_token":1.5e-7,"output_cost_per_token":6e-7,"cache_read_input_token_cost":3e-9}, + "output_cost_per_token": 4.2e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_vision": true, "supports_pdf_input": false, "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "input_cost_per_token": 4.62e-07, + "cache_read_input_token_cost": 8.8e-09, + "input_cost_per_token": 2.64e-07, "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.386e-06, + "output_cost_per_token": 7.92e-07, "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, "supports_function_calling": true, + "supports_pdf_input": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 1.54e-08, - "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"00:00-01:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"04:00-06:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"10:00-00:00"}],"input_cost_per_token":0.00000132,"output_cost_per_token":0.00000396,"cache_read_input_token_cost":4.4e-8}, - "supports_audio_input": false, - "supports_pdf_input": false, "supports_vision": false, "supports_web_search": false }, @@ -42755,14 +42930,14 @@ "openrouter/qwen/qwen3-coder-plus": { "cache_creation_input_token_cost": 8.125e-07, "cache_creation_input_token_cost_above_128k_tokens": 2.4375e-06, - "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, - "input_cost_per_token_above_32k_tokens": 1.17e-06, "cache_creation_input_token_cost_above_32k_tokens": 1.4625e-06, - "cache_read_input_token_cost_above_32k_tokens": 2.34e-07, - "output_cost_per_token_above_32k_tokens": 5.85e-06, "cache_read_input_token_cost": 1.3e-07, + "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, + "cache_read_input_token_cost_above_32k_tokens": 2.34e-07, + "deprecation_date": "2026-10-09", "input_cost_per_token": 6.5e-07, "input_cost_per_token_above_128k_tokens": 1.95e-06, + "input_cost_per_token_above_32k_tokens": 1.17e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, "max_output_tokens": 65536, @@ -42770,6 +42945,7 @@ "mode": "chat", "output_cost_per_token": 3.25e-06, "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 5.85e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42802,6 +42978,7 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-235b-a22b-thinking-2507": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2.3e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, @@ -43099,25 +43276,25 @@ "supports_web_search": false }, "openrouter/z-ai/glm-4.7": { - "input_cost_per_token": 4e-07, - "output_cost_per_token": 1.75e-06, "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 8e-08, + "cache_read_input_token_cost": 1.1e-07, + "input_cost_per_token": 6e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", + "output_cost_per_token": 2.2e-06, "source": "https://openrouter.ai/api/v1/models", - "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_vision": false, - "supports_prompt_caching": true, "supports_assistant_prefill": true, "supports_audio_input": false, + "supports_function_calling": true, "supports_pdf_input": false, + "supports_prompt_caching": true, + "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/z-ai/glm-4.7-flash": { @@ -43162,15 +43339,15 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.1": { - "input_cost_per_token": 9.66e-07, - "output_cost_per_token": 3.036e-06, - "cache_read_input_token_cost": 1.794e-07, "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 1.7914e-07, + "input_cost_per_token": 9.646e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", + "output_cost_per_token": 3.0316e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -56230,7 +56407,7 @@ "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", - "max_input_tokens": 1048576, + "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "realtime", @@ -56417,7 +56594,7 @@ "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", - "max_input_tokens": 1048576, + "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "realtime", @@ -59653,14 +59830,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 1e-06, + "cache_read_input_token_cost_batches": 5e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "prompt_cache_min_tokens": 512, "search_context_cost_per_query": { "search_context_size_high": 0.01, @@ -59693,14 +59874,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_batches": 1.25e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -60254,7 +60439,7 @@ "mode": "chat", "output_cost_per_token": 6.6e-07, "output_cost_per_token_priority": 8.25e-07, - "source": "https://api.fireworks.ai/v1/serverless/models", + "source": "https://api.fireworks.ai/v1/serverless/models?format=nested", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -60284,13 +60469,16 @@ }, "fireworks_ai/accounts/fireworks/models/deepseek-v4-flash-vision-exp": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "deprecation_date": "2026-09-25", "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, + "output_cost_per_token_priority": 8.25e-07, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, @@ -60353,7 +60541,7 @@ "mode": "chat", "output_cost_per_token": 6.6e-07, "output_cost_per_token_priority": 8.25e-07, - "source": "https://api.fireworks.ai/v1/serverless/models", + "source": "https://api.fireworks.ai/v1/serverless/models?format=nested", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -60383,13 +60571,16 @@ }, "fireworks_ai/deepseek-v4-flash-vision-exp": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "deprecation_date": "2026-09-25", "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, + "output_cost_per_token_priority": 8.25e-07, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, @@ -60518,14 +60709,17 @@ }, "fireworks_ai/muse-glimmer-30b": { "cache_read_input_token_cost": 4e-08, + "cache_read_input_token_cost_priority": 6e-08, "deprecation_date": "2026-09-25", "input_cost_per_token": 3.5e-07, + "input_cost_per_token_priority": 5.25e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, @@ -60567,14 +60761,17 @@ }, "fireworks_ai/accounts/fireworks/models/muse-glimmer-30b": { "cache_read_input_token_cost": 4e-08, + "cache_read_input_token_cost_priority": 6e-08, "deprecation_date": "2026-09-25", "input_cost_per_token": 3.5e-07, + "input_cost_per_token_priority": 5.25e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, @@ -62988,6 +63185,22 @@ "image" ] }, + "xai/grok-imagine-image-pro": { + "input_cost_per_image": 0.05, + "litellm_provider": "xai", + "mode": "image_generation", + "source": "https://docs.x.ai/docs/models", + "supported_endpoints": [ + "/v1/images/generations" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "image" + ] + }, "xai/grok-imagine-image-2.0": { "input_cost_per_image": 0.06, "litellm_provider": "xai", @@ -63994,6 +64207,62 @@ "supports_response_schema": true, "supports_vision": true }, + "azure_ai/deepseek-r1": { + "deprecation_date": "2026-08-13", + "input_cost_per_token": 1.35e-06, + "output_cost_per_token": 5.4e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/deepseek-v3-0324": { + "deprecation_date": "2026-07-13", + "input_cost_per_token": 1.14e-06, + "output_cost_per_token": 4.56e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/deepseek-v3.1": { + "deprecation_date": "2026-07-13", + "input_cost_per_token": 1.23e-06, + "output_cost_per_token": 4.94e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-3": { + "deprecation_date": "2026-05-01", + "input_cost_per_token": 3e-06, + "output_cost_per_token": 1.5e-05, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-3-mini": { + "deprecation_date": "2026-05-01", + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.27e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-4-fast-non-reasoning": { + "deprecation_date": "2026-05-01", + "input_cost_per_token": 2e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-4-fast-reasoning": { + "deprecation_date": "2026-05-01", + "input_cost_per_token": 2e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, "bedrock/us-gov-west-1/nvidia.nemotron-nano-3-30b": { "input_cost_per_token": 7.2e-08, "litellm_provider": "bedrock", @@ -64822,6 +65091,131 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "bedrock_mantle/deepseek.v3.1": { + "input_cost_per_token": 5.8e-07, + "output_cost_per_token": 1.68e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 128000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-1.html" + }, + "bedrock_mantle/moonshotai.kimi-k2-thinking": { + "input_cost_per_token": 6e-07, + "output_cost_per_token": 2.5e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html" + }, + "bedrock_mantle/qwen.qwen3-235b-a22b-2507": { + "input_cost_per_token": 2.2e-07, + "output_cost_per_token": 8.8e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-235b-a22b-2507.html" + }, + "bedrock_mantle/qwen.qwen3-32b": { + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 32000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-32b.html" + }, + "bedrock_mantle/qwen.qwen3-coder-30b-a3b-instruct": { + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-coder-30b-a3b-instruct.html" + }, + "bedrock_mantle/qwen.qwen3-coder-480b-a35b-instruct": { + "input_cost_per_token": 4.5e-07, + "output_cost_per_token": 1.8e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 128000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-coder-480b-a35b-instruct.html" + }, + "bedrock_mantle/qwen.qwen3-next-80b-a3b-instruct": { + "input_cost_per_token": 1.4e-07, + "output_cost_per_token": 1.2e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-next-80b-a3b.html" + }, + "bedrock_mantle/qwen.qwen3-vl-235b-a22b-instruct": { + "input_cost_per_token": 5.3e-07, + "output_cost_per_token": 2.66e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-vl-235b-a22b.html" + }, "azure/us-gov/gpt-5.1": { "cache_read_input_token_cost": 1.71875e-07, "default_reasoning_effort": "none", @@ -66172,23 +66566,23 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3-flash": { + "cache_read_input_token_cost": 1e-08, "input_cost_per_token": 4.5e-08, - "output_cost_per_token": 6e-07, - "cache_read_input_token_cost": 2.85e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", + "output_cost_per_token": 1.4e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-flash-vision-exp": { @@ -66349,24 +66743,24 @@ "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-flash-0731": { - "input_cost_per_token": 3e-08, - "output_cost_per_token": 3.2e-07, "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 2.2e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", + "output_cost_per_token": 3.2e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_response_schema": true, "supports_parallel_function_calling": true, "supports_pdf_input": false, - "supports_vision": false, "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/qwen/qwen3.7-flash": { @@ -66839,46 +67233,47 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.6-max-preview": { - "input_cost_per_token": 1.027e-06, - "output_cost_per_token": 6.162e-06, "cache_creation_input_token_cost": 1.28375e-06, - "input_cost_per_token_above_128k_tokens": 1.58e-06, - "output_cost_per_token_above_128k_tokens": 9.48e-06, "cache_creation_input_token_cost_above_128k_tokens": 1.975e-06, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 1.027e-06, + "input_cost_per_token_above_128k_tokens": 1.58e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 6.162e-06, + "output_cost_per_token_above_128k_tokens": 9.48e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, "openrouter/qwen/qwen3.6-27b": { - "input_cost_per_token": 3.2e-07, - "output_cost_per_token": 2.7e-06, "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_token": 3.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 262140, "max_tokens": 262140, "mode": "chat", + "output_cost_per_token": 3.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/openai/gpt-5.5-pro": { @@ -66923,23 +67318,23 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "input_cost_per_token": 4.9e-08, - "output_cost_per_token": 9.8e-08, - "cache_read_input_token_cost": 9.8e-09, + "cache_read_input_token_cost": 9.408e-09, + "input_cost_per_token": 4.704e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", + "output_cost_per_token": 9.408e-08, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/moonshotai/kimi-k2.6": { @@ -66964,22 +67359,22 @@ "supports_web_search": false }, "openrouter/google/gemma-4-26b-a4b-it": { - "cache_read_input_token_cost": 5e-08, - "input_cost_per_token": 9e-08, - "output_cost_per_token": 3e-07, + "cache_read_input_token_cost": 3.75e-08, + "input_cost_per_token": 6.75e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, "max_tokens": 235929, "mode": "chat", + "output_cost_per_token": 2.25e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": true, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67264,25 +67659,26 @@ "supports_video_input": true }, "openrouter/qwen/qwen3-max-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 7.8e-07, - "input_cost_per_token_above_32k_tokens": 1.56e-06, - "output_cost_per_token_above_32k_tokens": 7.8e-06, - "output_cost_per_token": 3.9e-06, "input_cost_per_token_above_128k_tokens": 1.95e-06, - "output_cost_per_token_above_128k_tokens": 9.75e-06, + "input_cost_per_token_above_32k_tokens": 1.56e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 3.9e-06, + "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 7.8e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -67534,59 +67930,62 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-vl-32b-instruct": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.04e-07, - "output_cost_per_token": 4.16e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4.16e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-8b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.8e-07, - "output_cost_per_token": 2.1e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.1e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-8b-instruct": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.17e-07, - "output_cost_per_token": 4.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4.55e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67616,40 +68015,41 @@ "supports_web_search": true }, "openrouter/qwen/qwen3-vl-30b-a3b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2e-07, - "output_cost_per_token": 2.4e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-30b-a3b-instruct": { - "input_cost_per_token": 1.3e-07, - "output_cost_per_token": 5.2e-07, + "input_cost_per_token": 1.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 6e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67673,21 +68073,22 @@ "supports_web_search": true }, "openrouter/qwen/qwen3-vl-235b-a22b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 4e-07, - "output_cost_per_token": 4e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67712,32 +68113,33 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-max": { - "input_cost_per_token": 7.8e-07, - "input_cost_per_token_above_32k_tokens": 1.56e-06, - "cache_creation_input_token_cost_above_32k_tokens": 1.95e-06, - "cache_read_input_token_cost_above_32k_tokens": 3.12e-07, - "output_cost_per_token_above_32k_tokens": 7.8e-06, - "output_cost_per_token": 3.9e-06, - "cache_read_input_token_cost": 1.56e-07, "cache_creation_input_token_cost": 9.75e-07, - "input_cost_per_token_above_128k_tokens": 1.95e-06, - "output_cost_per_token_above_128k_tokens": 9.75e-06, - "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, "cache_creation_input_token_cost_above_128k_tokens": 2.4375e-06, + "cache_creation_input_token_cost_above_32k_tokens": 1.95e-06, + "cache_read_input_token_cost": 1.56e-07, + "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, + "cache_read_input_token_cost_above_32k_tokens": 3.12e-07, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 7.8e-07, + "input_cost_per_token_above_128k_tokens": 1.95e-06, + "input_cost_per_token_above_32k_tokens": 1.56e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 3.9e-06, + "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 7.8e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, - "supports_response_schema": true, - "supports_vision": false, "supports_prompt_caching": true, "supports_reasoning": false, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/deepseek/deepseek-v3.1-terminus": { @@ -67832,23 +68234,24 @@ "openrouter/qwen/qwen-plus-2025-07-28": { "cache_creation_input_token_cost": 3.25e-07, "cache_read_input_token_cost": 5.2e-08, + "deprecation_date": "2026-10-09", "input_cost_per_token": 2.6e-07, - "output_cost_per_token": 7.8e-07, "input_cost_per_token_above_256k_tokens": 7.8e-07, - "output_cost_per_token_above_256k_tokens": 2.34e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 7.8e-07, + "output_cost_per_token_above_256k_tokens": 2.34e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -67872,21 +68275,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-30b-a3b-thinking-2507": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2e-07, - "output_cost_per_token": 2.4e-06, "litellm_provider": "openrouter", "max_input_tokens": 81920, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -68195,21 +68599,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-8b": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.17e-07, - "output_cost_per_token": 4.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", + "output_cost_per_token": 4.55e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -68252,21 +68657,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-235b-a22b": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 4.55e-07, - "output_cost_per_token": 1.82e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", + "output_cost_per_token": 1.82e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -71335,6 +71741,23 @@ "output_cost_per_token": 0.0, "source": "https://openrouter.ai/typesafe/jev-1.13" }, + "openrouter/typesafe/jev-router": { + "input_cost_per_token": 0, + "output_cost_per_token": 0, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "source": "https://openrouter.ai/typesafe/jev-router", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_audio_input": true, + "supports_video_input": true + }, "typesafe/jev-1.13.0": { "input_cost_per_token": 4.2e-08, "litellm_provider": "typesafe", @@ -73019,14 +73442,14 @@ "supports_web_search": false }, "openrouter/inclusionai/ling-3.0-flash-vl": { - "cache_read_input_token_cost": 1.2e-08, - "input_cost_per_token": 6e-08, + "cache_read_input_token_cost": 4.2e-09, + "input_cost_per_token": 2.1e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 1.8e-07, + "output_cost_per_token": 6.16e-08, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -76537,6 +76960,23 @@ "supports_vision": true, "supports_web_search": false }, + "openrouter/perceptron/perceptron-mk1.5": { + "input_cost_per_token": 1.5e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 36864, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true + }, "vertex_ai/gemini-2.0-flash": { "deprecation_date": "2026-06-01", "input_cost_per_audio_token": 1e-06, diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index fa1828c780a..c4048cac905 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -103,6 +103,11 @@ "minimum": 0, "description": "Rate applied once the prompt exceeds the token threshold in the field name." }, + "cache_creation_input_token_cost_above_200k_tokens_batches": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "cache_creation_input_token_cost_above_256k_tokens": { "type": "number", "minimum": 0, diff --git a/tests/README.MD b/tests/README.MD index 57275a031f7..6a5da137203 100644 --- a/tests/README.MD +++ b/tests/README.MD @@ -4,6 +4,6 @@ To make it easier to contribute and map what behavior is tested, -we've started mapping the litellm directory in `tests/test_litellm` +we've started mapping the litellm directory in `tests/unit` This folder can only run mock tests. diff --git a/tests/e2e/batches/batch_cleanup.py b/tests/e2e/batches/batch_cleanup.py index 86e47c0b1e1..f889844f1ae 100644 --- a/tests/e2e/batches/batch_cleanup.py +++ b/tests/e2e/batches/batch_cleanup.py @@ -121,7 +121,8 @@ def cleanup_batch( if current.status == "cancelling" and not needs_terminal_state: return assert clock() < deadline, ( - f"Batch {batch_id} cancellation did not finish within {BATCH_CANCEL_TIMEOUT_SECONDS}s" + f"Batch {batch_id} cancellation did not finish within {BATCH_CANCEL_TIMEOUT_SECONDS}s, " + f"last status {current.status}" ) wait(BATCH_CANCEL_POLL_SECONDS) diff --git a/tests/e2e/batches/bedrock_env_gateway.py b/tests/e2e/batches/bedrock_env_gateway.py index fb3ec60c87c..0b1d840eb30 100644 --- a/tests/e2e/batches/bedrock_env_gateway.py +++ b/tests/e2e/batches/bedrock_env_gateway.py @@ -8,6 +8,7 @@ a batch create through it proves blank means unset, not an empty string. from __future__ import annotations +import importlib.util import os import shutil import socket @@ -28,7 +29,13 @@ from pydantic import TypeAdapter STARTUP_TIMEOUT_SECONDS: Final = 240 LOG_TAIL_BYTES: Final = 4000 -REPO_ROOT: Final = Path(__file__).resolve().parents[3] + + +def litellm_root() -> Path: + spec: Final = importlib.util.find_spec("litellm") + assert spec is not None and spec.origin is not None, "litellm must be importable to boot the blank-S3-env gateway" + return Path(spec.origin).resolve().parents[1] + _CONFIG_YAML: Final = """model_list: - model_name: bedrock-blank-s3-batch @@ -68,6 +75,7 @@ class BedrockEnvGateway: @classmethod def start(cls) -> BedrockEnvGateway: assert os.environ.get("DATABASE_URL"), "DATABASE_URL is required for the blank-S3-env gateway" + root: Final = litellm_root() port: Final = available_port() base_url: Final = f"http://127.0.0.1:{port}" master_key: Final = f"sk-e2e-blank-s3-{unique_marker()}" @@ -79,7 +87,7 @@ class BedrockEnvGateway: "DATABASE_URL": os.environ["DATABASE_URL"], "LITELLM_MASTER_KEY": master_key, "STORE_MODEL_IN_DB": "False", - "PYTHONPATH": str(REPO_ROOT), + "PYTHONPATH": str(root), "AWS_S3_ENCRYPTION_KEY_ID": "", "AWS_S3_BUCKET_OWNER": "", } @@ -113,13 +121,11 @@ class BedrockEnvGateway: stdout=log, stderr=log, start_new_session=True, - cwd=REPO_ROOT, + cwd=root, ) deadline: Final = time.monotonic() + STARTUP_TIMEOUT_SECONDS while time.monotonic() < deadline: - assert gateway._child.poll() is None, ( - f"blank-S3-env gateway exited early; log tail:\n{gateway.log_tail()}" - ) + assert gateway._child.poll() is None, f"blank-S3-env gateway exited early; log tail:\n{gateway.log_tail()}" result = gateway.proxy.transport.probe("/health/liveliness", params=NoBody()) if result.status_code == 200: return gateway diff --git a/tests/e2e/batches/test_batch_cleanup.py b/tests/e2e/batches/test_batch_cleanup.py index a0932a80dfe..28eb362e876 100644 --- a/tests/e2e/batches/test_batch_cleanup.py +++ b/tests/e2e/batches/test_batch_cleanup.py @@ -210,6 +210,7 @@ class TestBatchCancellation: with pytest.raises(ExceptionGroup) as caught: manager.teardown() assert "cancellation did not finish" in str(caught.value.exceptions[0]) + assert "last status cancelling" in str(caught.value.exceptions[0]) client.calls.assert_done() @pytest.mark.parametrize("status", ["completed", "failed", "expired", "cancelled"]) diff --git a/tests/e2e/logging/test_otel_trace_e2e.py b/tests/e2e/logging/test_otel_trace_e2e.py index 8d154ca0837..8d8595cf221 100644 --- a/tests/e2e/logging/test_otel_trace_e2e.py +++ b/tests/e2e/logging/test_otel_trace_e2e.py @@ -2,8 +2,9 @@ Covers logging.otel.success.exports_metric: a successful non-streaming call must land at the OTEL destination as ONE connected trace - a single root SERVER span -with the auth phase, db lookups, and cost write under it, and the gen-AI CLIENT -span parented into the same tree. The regression this pins: the proxy publishing +with the auth phase and db lookups under it, the gen-AI CLIENT span parented +into the same tree, and the cost write either under it or as the root of its +own trace linked back to the request span. The regression this pins: the proxy publishing the global TracerProvider before callbacks init made server spans export through a different provider than the preset's gen-AI spans, so the destination received the gen-AI span alone, dangling (fixed in #30590; verified failing at its parent @@ -28,7 +29,7 @@ from e2e_config import CHEAP_ANTHROPIC_MODEL, CHEAP_OPENAI_MODEL, OTEL_EXPORTER_ from lifecycle import ResourceManager from logging_client import INVALID_UPSTREAM_API_KEY, LoggingClient, first_ok, readiness_details_body from models import LiteLLMParamsBody -from otel_client import JaegerSpan, JaegerTrace, OtelReader +from otel_client import CallTraces, JaegerSpan, JaegerTrace, OtelReader, root_span from pydantic import BaseModel, ConfigDict, ValidationError pytestmark = pytest.mark.e2e @@ -78,12 +79,13 @@ def _chain_reaches(span_id: str, root_id: str, trace: JaegerTrace) -> bool: return False -def _assert_complete_trace( - hits: list[JaegerTrace], *, route: str, genai_span: str, require_cost_span: bool = True -) -> None: - """The enforced behavior: the destination holds exactly one trace for the - call, rooted at the SERVER span, with auth/db/cost children and the gen-AI - span all connected into that one tree - no dangling parent references.""" +def _assert_complete_trace(traces: CallTraces, *, route: str, genai_span: str, require_cost_span: bool = True) -> None: + """The enforced behavior: the destination holds exactly one call-id-tagged + trace for the call, rooted at the SERVER span, with auth/db children and + the gen-AI span all connected into that one tree - no dangling parent + references - and the cost write either in that trace or as the root of + its own trace linked FOLLOWS_FROM to the request SERVER span.""" + hits = traces.hits assert hits, ( "no trace for this call arrived at the destination within the deadline " "(nothing tagged with its call id was found)" @@ -119,8 +121,20 @@ def _assert_complete_trace( assert any(name.startswith(DB_SPAN_PREFIX) for name in names), ( f"no db ('{DB_SPAN_PREFIX}*') span in the trace; spans: {names}" ) - if require_cost_span: - assert COST_SPAN in names, f"cost write span {COST_SPAN!r} missing; spans: {names}" + if require_cost_span and COST_SPAN not in names: + cost_traces = [t for t in traces.linked if (r := root_span(t)) is not None and r.operation_name == COST_SPAN] + assert len(cost_traces) == 1, ( + f"cost write span {COST_SPAN!r} reached neither the request trace nor its own " + f"trace linked to the request SERVER span; request spans: {names}; " + f"linked traces: {[(t.trace_id, t.span_names()) for t in traces.linked]}" + ) + cost_root = root_span(cost_traces[0]) + assert cost_root is not None, f"cost write trace has no single root; spans: {cost_traces[0].span_names()}" + link = next(ref for ref in cost_root.references if ref.span_id == root.span_id) + assert link.ref_type == "FOLLOWS_FROM" and link.trace_id == trace.trace_id, ( + f"the cost write trace's root must reference the request SERVER span FOLLOWS_FROM, " + f"got refType={link.ref_type!r} traceID={link.trace_id!r} (request trace {trace.trace_id})" + ) genai = next((span for span in trace.spans if span.operation_name == genai_span), None) assert genai is not None, f"gen-AI span {genai_span!r} missing; spans: {names}" @@ -131,9 +145,15 @@ def _assert_complete_trace( ) -def _settled_names(*, route: str, genai_span: str, require_cost_span: bool = True) -> set[str]: - names = {f"POST {route}", f"auth {route}", genai_span} - return (names | {COST_SPAN}) if require_cost_span else names +def _poll( + otel_reader: OtelReader, *, call_id: str, route: str, genai_span: str, require_cost_span: bool = True +) -> CallTraces: + return otel_reader.poll_traces_for_call( + call_id=call_id, + settled_names={f"POST {route}", f"auth {route}", genai_span}, + settled_prefixes={DB_SPAN_PREFIX}, + linked_names=frozenset({COST_SPAN}) if require_cost_span else frozenset(), + ) def _tag(span: JaegerSpan, key: str) -> str | int | float | bool | None: @@ -174,7 +194,7 @@ def one_served_genai_span(trace: JaegerTrace, genai_span: str) -> JaegerSpan: return served[0] -def _assert_real_ttft(hits: list[JaegerTrace], *, genai_span: str) -> None: +def _assert_real_ttft(hits: tuple[JaegerTrace, ...], *, genai_span: str) -> None: """The enforced behavior: the gen-AI span for the attempt that served the stream records a TTFT that is a real measurement - present, numeric, positive, and strictly less than that span's own total duration. A TTFT of @@ -286,9 +306,10 @@ class TestOtelTraceCompleteness: /chat/completions request produces one complete OTEL trace. The trace should have a single server root span for the incoming request, with - the authentication, database, and cost-recording work beneath it. The span for - the actual model call must also belong to that same trace, rather than being - exported separately with a missing parent. + the authentication and database work beneath it. The span for the actual model + call must also belong to that same trace, rather than being exported separately + with a missing parent, and the cost-recording work must land either in that + trace or in its own trace linked to it. This matters because a split trace is easy to miss: all of the spans may still arrive, but the model call appears without the surrounding request context. @@ -308,12 +329,8 @@ class TestOtelTraceCompleteness: outcome = first_ok(client, lambda: client.chat_raw(key, MODEL, f"reply with one word {marker}", max_tokens=16)) assert outcome.call_id is not None, "success response must carry x-litellm-call-id" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=f"chat {MODEL}"), - settled_prefixes={DB_SPAN_PREFIX}, - ) - _assert_complete_trace(hits, route=route, genai_span=f"chat {MODEL}") + traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=f"chat {MODEL}") + _assert_complete_trace(traces, route=route, genai_span=f"chat {MODEL}") @pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["chat_completions"]) @pytest.mark.otel_tls @@ -337,11 +354,7 @@ class TestOtelTraceCompleteness: ) assert outcome.call_id is not None, "success response must carry x-litellm-call-id" - hits: Final = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=f"chat {MODEL}"), - settled_prefixes={DB_SPAN_PREFIX}, - ) + hits: Final = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=f"chat {MODEL}") _assert_complete_trace(hits, route=route, genai_span=f"chat {MODEL}") @pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["messages"]) @@ -352,8 +365,10 @@ class TestOtelTraceCompleteness: produces exactly one complete OTEL trace. The trace must have a single root span named "POST /v1/messages". The - authentication, database, cost-writing, and model-call spans must all belong to + authentication, database, and model-call spans must all belong to the same trace and have valid parent relationships leading back to that root. + The cost-writing span must land in the request trace or in its own trace + linked to it. The model-call span is expected to be named "chat ". The test fails if the request is split across multiple traces, if any span references a missing @@ -370,12 +385,8 @@ class TestOtelTraceCompleteness: ) assert outcome.call_id is not None, "success response must carry x-litellm-call-id" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=f"chat {MODEL}"), - settled_prefixes={DB_SPAN_PREFIX}, - ) - _assert_complete_trace(hits, route=route, genai_span=f"chat {MODEL}") + traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=f"chat {MODEL}") + _assert_complete_trace(traces, route=route, genai_span=f"chat {MODEL}") @pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["responses"]) def test_responses_exports_complete_trace( @@ -385,12 +396,14 @@ class TestOtelTraceCompleteness: produces exactly one complete OTEL trace. The trace must have a single root span named "POST /v1/responses". The - authentication, database, cost-writing, and model-call spans must all belong to - the same trace and have valid parent relationships leading back to that root. + authentication, database, and model-call spans must all belong to the same + trace and have valid parent relationships leading back to that root. The cost + write finishes after the response, so it lands as the root of its own trace + linked FOLLOWS_FROM to the request SERVER span. - The model-call span is expected to be named "chat ". The test fails if - the request is split across multiple traces, if any span references a missing - parent, or if the model-call span cannot be connected back to the root.""" + The model-call span is expected to be named "chat ". The test fails on + a split request trace, a dangling parent, a disconnected model-call span, or + a cost write that is neither in the request trace nor linked to it.""" route = "/v1/responses" _assert_otel_destination_configured(client) @@ -405,12 +418,8 @@ class TestOtelTraceCompleteness: assert outcome.call_id is not None, "success response must carry x-litellm-call-id" genai_span = f"chat {CHEAP_OPENAI_MODEL}" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=genai_span), - settled_prefixes={DB_SPAN_PREFIX}, - ) - _assert_complete_trace(hits, route=route, genai_span=genai_span) + traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span) + _assert_complete_trace(traces, route=route, genai_span=genai_span) @pytest.mark.covers("logging.otel.stream.exports_metric", exercised_on=["chat_completions"]) def test_chat_completions_stream_exports_complete_trace( @@ -418,8 +427,9 @@ class TestOtelTraceCompleteness: ) -> None: """A successful streamed `/chat/completions` request should export one complete OTEL trace. The trace must contain a single root `SERVER` - span, with the auth, database, cost, and gen-AI `CLIENT` spans all - connected back to that root. + span, with the auth, database, and gen-AI `CLIENT` spans all + connected back to that root, and the cost write in that trace or in + its own trace linked to it. Streaming has an additional lifecycle risk because the gen-AI span is closed by the stream-consumption path after the final chunk has @@ -451,14 +461,10 @@ class TestOtelTraceCompleteness: ) genai_span = f"chat {MODEL}" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=genai_span), - settled_prefixes={DB_SPAN_PREFIX}, - ) - _assert_complete_trace(hits, route=route, genai_span=genai_span) + traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span) + _assert_complete_trace(traces, route=route, genai_span=genai_span) - served = one_served_genai_span(hits[0], genai_span) + served = one_served_genai_span(traces.hits[0], genai_span) assert _tag(served, "litellm.request.streaming") is True, ( "the gen-AI span must record litellm.request.streaming=true; its absence means " "the stream flag was dropped before the model call" @@ -470,8 +476,9 @@ class TestOtelTraceCompleteness: ) -> None: """A successful streamed `/v1/messages` request should export one complete OTEL trace. The trace must contain a single root `SERVER` - span, with the auth, database, cost, and gen-AI `CLIENT` spans all - connected back to that root. + span, with the auth, database, and gen-AI `CLIENT` spans all + connected back to that root, and the cost write in that trace or in + its own trace linked to it. This endpoint has the same streaming lifecycle risk as `/chat/completions`: the gen-AI span is closed by the @@ -503,14 +510,10 @@ class TestOtelTraceCompleteness: ) genai_span = f"chat {MODEL}" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=genai_span), - settled_prefixes={DB_SPAN_PREFIX}, - ) - _assert_complete_trace(hits, route=route, genai_span=genai_span) + traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span) + _assert_complete_trace(traces, route=route, genai_span=genai_span) - served = one_served_genai_span(hits[0], genai_span) + served = one_served_genai_span(traces.hits[0], genai_span) assert _tag(served, "litellm.request.streaming") is True, ( "the gen-AI span must record litellm.request.streaming=true; its absence means " "the stream flag was dropped before the model call" @@ -555,14 +558,12 @@ class TestOtelTraceCompleteness: ) genai_span = f"chat {CHEAP_OPENAI_MODEL}" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=genai_span, require_cost_span=False), - settled_prefixes={DB_SPAN_PREFIX}, + traces = _poll( + otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False ) - _assert_complete_trace(hits, route=route, genai_span=genai_span, require_cost_span=False) + _assert_complete_trace(traces, route=route, genai_span=genai_span, require_cost_span=False) - one_served_genai_span(hits[0], genai_span) + one_served_genai_span(traces.hits[0], genai_span) spend_row = client.poll_proxy_spend_for_key(key) assert spend_row is not None and spend_row.spend is not None and spend_row.spend > 0, ( @@ -609,12 +610,8 @@ class TestOtelTraceCompleteness: ) genai_span = f"chat {MODEL}" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=genai_span), - settled_prefixes={DB_SPAN_PREFIX}, - ) - _assert_real_ttft(hits, genai_span=genai_span) + traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span) + _assert_real_ttft(traces.hits, genai_span=genai_span) @pytest.mark.covers("logging.otel.stream.records_ttft", exercised_on=["messages"]) def test_messages_stream_records_real_ttft( @@ -651,12 +648,8 @@ class TestOtelTraceCompleteness: ) genai_span = f"chat {MODEL}" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=genai_span), - settled_prefixes={DB_SPAN_PREFIX}, - ) - _assert_real_ttft(hits, genai_span=genai_span) + traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span) + _assert_real_ttft(traces.hits, genai_span=genai_span) @pytest.mark.covers("logging.otel.stream.records_ttft", exercised_on=["responses"]) def test_responses_stream_records_real_ttft( @@ -693,12 +686,10 @@ class TestOtelTraceCompleteness: ) genai_span = f"chat {CHEAP_OPENAI_MODEL}" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=genai_span, require_cost_span=False), - settled_prefixes={DB_SPAN_PREFIX}, + traces = _poll( + otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False ) - _assert_real_ttft(hits, genai_span=genai_span) + _assert_real_ttft(traces.hits, genai_span=genai_span) @pytest.mark.covers("logging.otel.failure.exports_metric", exercised_on=["chat_completions"]) def test_failed_chat_completions_error_span_attributes( @@ -745,18 +736,16 @@ class TestOtelTraceCompleteness: assert outcome.call_id is not None, "failed responses must still carry x-litellm-call-id" genai_span = f"chat {model_name}" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=genai_span, require_cost_span=False), - settled_prefixes={DB_SPAN_PREFIX}, + traces = _poll( + otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False ) - _assert_complete_trace(hits, route=route, genai_span=genai_span, require_cost_span=False) + _assert_complete_trace(traces, route=route, genai_span=genai_span, require_cost_span=False) - root = next(span for span in hits[0].spans if not span.references) + root = next(span for span in traces.hits[0].spans if not span.references) assert str(_tag(root, "http.status_code")) == "401", ( f"the SERVER span must record the 401 the client received, got {_tag(root, 'http.status_code')!r}" ) - genai = next(span for span in hits[0].spans if span.operation_name == genai_span) + genai = next(span for span in traces.hits[0].spans if span.operation_name == genai_span) _assert_error_span_contract(genai) @pytest.mark.covers("logging.otel.failure.exports_metric", exercised_on=["messages"]) @@ -803,16 +792,14 @@ class TestOtelTraceCompleteness: assert outcome.call_id is not None, "failed responses must still carry x-litellm-call-id" genai_span = f"chat {model_name}" - hits = otel_reader.poll_traces_for_call( - call_id=outcome.call_id, - settled_names=_settled_names(route=route, genai_span=genai_span, require_cost_span=False), - settled_prefixes={DB_SPAN_PREFIX}, + traces = _poll( + otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False ) - _assert_complete_trace(hits, route=route, genai_span=genai_span, require_cost_span=False) + _assert_complete_trace(traces, route=route, genai_span=genai_span, require_cost_span=False) - root = next(span for span in hits[0].spans if not span.references) + root = next(span for span in traces.hits[0].spans if not span.references) assert str(_tag(root, "http.status_code")) == "401", ( f"the SERVER span must record the 401 the client received, got {_tag(root, 'http.status_code')!r}" ) - genai = next(span for span in hits[0].spans if span.operation_name == genai_span) + genai = next(span for span in traces.hits[0].spans if span.operation_name == genai_span) _assert_error_span_contract(genai) diff --git a/tests/e2e/otel_client.py b/tests/e2e/otel_client.py index b11fddebc9c..c5d709048b3 100644 --- a/tests/e2e/otel_client.py +++ b/tests/e2e/otel_client.py @@ -10,6 +10,11 @@ so the completeness assertions see the whole tree. A failed query is a hard failure, never an empty result - an unreachable destination must not read as "the trace never arrived". +Service spans that end after the response (the cost write is one) carry no +call id and land as the root of their own trace with a link back to the +request span, so they are fetched by operation name and matched by that link +to the request root rather than by the tag query. + External reads go through ``e2e_http`` (the only module allowed to call ``requests.*``). """ @@ -18,7 +23,9 @@ from __future__ import annotations import json import time +from collections.abc import Iterator from dataclasses import dataclass +from typing import Final import pytest from pydantic import BaseModel, ConfigDict, Field @@ -84,18 +91,57 @@ class JaegerTracesPage(BaseModel): class _TracesQuery(BaseModel): service: str - tags: str + tags: str | None = None + operation: str | None = None limit: int = 20 lookback: str = "1h" + start: int | None = None + end: int | None = None + + +def _ticks() -> Iterator[None]: + while True: + yield None + time.sleep(POLL_INTERVAL) def _settled(trace: JaegerTrace, names: set[str], prefixes: set[str]) -> bool: present = set(trace.span_names()) - return names.issubset(present) and all( - any(name.startswith(prefix) for name in present) for prefix in prefixes + return names.issubset(present) and all(any(name.startswith(prefix) for name in present) for prefix in prefixes) + + +def root_span(trace: JaegerTrace) -> JaegerSpan | None: + """The single span whose references all point outside the trace (a span + with no references qualifies). None when there is not exactly one.""" + in_trace = {span.span_id for span in trace.spans} + roots = [span for span in trace.spans if all(ref.span_id not in in_trace for ref in span.references)] + return roots[0] if len(roots) == 1 else None + + +def _follows(trace: JaegerTrace, parent_trace_id: str, parent_span_id: str) -> bool: + root = root_span(trace) + return root is not None and any( + ref.trace_id == parent_trace_id and ref.span_id == parent_span_id for ref in root.references ) +@dataclass(frozen=True, slots=True) +class CallTraces: + hits: tuple[JaegerTrace, ...] + linked: tuple[JaegerTrace, ...] + + +@dataclass(frozen=True, slots=True) +class _Observation: + traces: CallTraces + missing: tuple[str, ...] + unreachable: NetworkError | None + + def settled(self, names: set[str], prefixes: set[str]) -> bool: + hits: Final = self.traces.hits + return self.unreachable is None and len(hits) == 1 and not self.missing and _settled(hits[0], names, prefixes) + + @dataclass(frozen=True, slots=True) class OtelReader: query_url: str @@ -119,36 +165,95 @@ class OtelReader: case failure: pytest.fail(f"Jaeger query API at {self.query_url} failed: {failure}") + def _query_operation(self, operation: str, *, start: int) -> Result[JaegerTracesPage]: + return get( + URL(f"{self.query_url}/api/traces"), + headers=NoBody(), + params=_TracesQuery( + service=JAEGER_SERVICE, + operation=operation, + limit=200, + start=start, + end=int(time.time() * 1_000_000), + ), + response_type=JaegerTracesPage, + timeout=30.0, + ) + + def linked_traces(self, *, operation: str, parent: JaegerTrace) -> tuple[JaegerTrace, ...] | NetworkError: + """Traces whose root span references the parent trace's root span. + Detached post-response work lands as the root of its own trace with a + link back to the request span instead of the call-id tag, so it is + found by operation name, windowed to start at the parent root's start + time (the detached span always starts after it), and matched on that + link. A NetworkError is handed back so the polling caller can tell an + unreachable read-back endpoint from a span that never arrived.""" + parent_root: Final = root_span(parent) + if parent_root is None: + return () + match self._query_operation(operation, start=parent_root.start_time): + case Success(data=page): + return tuple(t for t in page.data if _follows(t, parent.trace_id, parent_root.span_id)) + case NetworkError() as failure: + return failure + case failure: + pytest.fail(f"Jaeger query API at {self.query_url} failed: {failure}") + def poll_traces_for_call( - self, *, call_id: str, settled_names: set[str], settled_prefixes: set[str] - ) -> list[JaegerTrace]: - """Poll until exactly one trace holds the call and it carries every span + self, + *, + call_id: str, + settled_names: set[str], + settled_prefixes: set[str], + linked_names: frozenset[str] = frozenset(), + ) -> CallTraces: + """Poll until exactly one trace holds the call, it carries every span name in ``settled_names`` plus at least one name per prefix in - ``settled_prefixes`` (spans flush in batches, the cost write lands after - the response), then return the hits. At the deadline the last hits are - returned as-is so the caller's assertions report the real final state - - on a split trace this never settles and the orphan comes back.""" - deadline = time.monotonic() + POLL_TIMEOUT - hits: list[JaegerTrace] = [] - unreachable: NetworkError | None = None - while time.monotonic() < deadline: - match self._query_traces(call_id): - case Success(data=page): - unreachable = None - hits = page.data - if len(hits) == 1 and _settled(hits[0], settled_names, settled_prefixes): - return hits - case NetworkError() as failure: - unreachable = failure - case failure: - pytest.fail(f"Jaeger query API at {self.query_url} failed: {failure}") - time.sleep(POLL_INTERVAL) - if unreachable is not None: + ``settled_prefixes``, and every name in ``linked_names`` is either in + that trace or is the root of its own trace referencing the request + root (post-response work detaches per #42826). At the deadline the + last observed state is returned as-is so the caller's assertions + report the real final state - on a split trace this never settles and + the orphan comes back. A read-back endpoint still failing at the + deadline (either query) is a hard failure, not a missing span.""" + deadline: Final = time.monotonic() + POLL_TIMEOUT + last: Final = self._poll(call_id, settled_names, settled_prefixes, linked_names, deadline) + if last.unreachable is not None: pytest.fail( f"Jaeger query API at {self.query_url} stayed unreachable until the " - f"{POLL_TIMEOUT}s poll deadline: {unreachable}" + f"{POLL_TIMEOUT}s poll deadline: {last.unreachable}" ) - return hits + return last.traces + + def _observe(self, call_id: str, linked_names: frozenset[str]) -> _Observation: + match self._query_traces(call_id): + case NetworkError() as failure: + return _Observation(CallTraces((), ()), tuple(linked_names), failure) + case Success(data=page): + if len(page.data) != 1: + return _Observation(CallTraces(tuple(page.data), ()), tuple(linked_names), None) + hit: Final = page.data[0] + present: Final = frozenset(hit.span_names()) + results: Final = { + name: self.linked_traces(operation=name, parent=hit) for name in linked_names if name not in present + } + unreachable: Final = next((r for r in results.values() if isinstance(r, NetworkError)), None) + linked: Final = tuple(t for r in results.values() if not isinstance(r, NetworkError) for t in r) + missing: Final = tuple(name for name, r in results.items() if isinstance(r, NetworkError) or not r) + return _Observation(CallTraces((hit,), linked), missing, unreachable) + case failure: + pytest.fail(f"Jaeger query API at {self.query_url} failed: {failure}") + + def _poll( + self, + call_id: str, + names: set[str], + prefixes: set[str], + linked_names: frozenset[str], + deadline: float, + ) -> _Observation: + observations: Final = (self._observe(call_id, linked_names) for _ in _ticks()) + return next(o for o in observations if o.settled(names, prefixes) or time.monotonic() >= deadline) def build_otel_reader() -> OtelReader: diff --git a/tests/e2e/quota_management/spend_tracking/test_spend_routes.py b/tests/e2e/quota_management/spend_tracking/test_spend_routes.py index 67fd88bc84d..c3697a31424 100644 --- a/tests/e2e/quota_management/spend_tracking/test_spend_routes.py +++ b/tests/e2e/quota_management/spend_tracking/test_spend_routes.py @@ -74,6 +74,8 @@ SPEND_ROUTES = ( _SPEND_PREFIXES = ("/spend", "/global/spend", "/global/activity") +_CAPTURE_RATE_ROUTE: Final = "/spend/capture_rate" + # Served from the MonthlyGlobalSpend / DailyTagSpend / Last30d* views, which the # proxy creates in the background once the schema migrations have landed, so on a # fresh database they can 500 for a while after the proxy starts serving. @@ -120,7 +122,7 @@ def test_schema_listed_spend_routes_are_responsive(client: SpendClient) -> None: and "{" not in path and any(path.startswith(prefix) for prefix in _SPEND_PREFIXES) ] - extras = [path for path in discovered if path not in SPEND_ROUTES] + extras = [path for path in discovered if path not in (*SPEND_ROUTES, _CAPTURE_RATE_ROUTE)] params = _date_range() results = [(path, client.probe(path, params=params)) for path in extras] @@ -132,3 +134,11 @@ def test_schema_listed_spend_routes_are_responsive(client: SpendClient) -> None: if not result.healthy ] assert not offenders, "non-responsive schema spend routes:\n" + "\n".join(offenders) + + +def test_capture_rate_reports_or_names_the_missing_billing_key(client: SpendClient) -> None: + result: Final = client.probe(_CAPTURE_RATE_ROUTE, params=_date_range()) + print(f"{_CAPTURE_RATE_ROUTE} -> {result.status_code}\n{result.body[:600]}") + assert result.status_code == 200 or (result.status_code == 503 and "OPENAI_ADMIN_KEY is not set" in result.body), ( + f"{_CAPTURE_RATE_ROUTE} -> {result.status_code}\n{result.body[:600]}" + ) diff --git a/tests/e2e/quota_management/spend_tracking/test_websearch_interception_session_e2e.py b/tests/e2e/quota_management/spend_tracking/test_websearch_interception_session_e2e.py index 89d0beec414..6352ab67c3c 100644 --- a/tests/e2e/quota_management/spend_tracking/test_websearch_interception_session_e2e.py +++ b/tests/e2e/quota_management/spend_tracking/test_websearch_interception_session_e2e.py @@ -12,13 +12,14 @@ Needs a proxy booted with the callback and a real search backend, which ``gateway/stage_mirror_ci_config.yml`` carries as the ``e2e-search`` Perplexity tool. """ -from typing import Final +from typing import Final, Literal import pytest from e2e_config import unique_marker from e2e_http import unwrap from lifecycle import ResourceManager from models import ( + AnthropicContentBlock, AnthropicMessagesBody, AnthropicWebSearchTool, ChatMessage, @@ -26,6 +27,7 @@ from models import ( SpendLogRow, ) from proxy_client import ProxyClient +from pydantic import BaseModel, ValidationError pytestmark = pytest.mark.e2e @@ -37,6 +39,18 @@ def _has_search_row(rows: list[SpendLogRow]) -> bool: return any(row.call_type == SEARCH_CALL_TYPE for row in rows) +class _SearchResultError(BaseModel): + type: Literal["web_search_tool_result_error"] + error_code: str + + +def _search_error_code(block: AnthropicContentBlock) -> str | None: + try: + return _SearchResultError.model_validate((block.model_extra or {}).get("content")).error_code + except ValidationError: + return None + + class TestWebSearchInterceptionSession: @pytest.mark.covers( "quota_management.spend_tracking.websearch_interception.bills_under_request_session", @@ -79,6 +93,15 @@ class TestWebSearchInterceptionSession: f"precondition: the turn never ran an intercepted search, so there is no search row to attribute. " f"blocks={block_types}" ) + search_errors: Final = tuple( + code + for block in response.content or () + if block.type == "web_search_tool_result" and (code := _search_error_code(block)) is not None + ) + assert not search_errors, ( + f"precondition: the e2e-search tool failed upstream ({search_errors}), so no {SEARCH_CALL_TYPE} row is " + "billed at all; check the proxy's search tool credentials before reading this as a session bug" + ) rows: Final = proxy.poll_logs_for_session(session_id, min_rows=2, predicate=_has_search_row) by_call_type: Final = {row.call_type or "" for row in rows} diff --git a/tests/e2e/secret_manager/secret_store_cyberark.py b/tests/e2e/secret_manager/secret_store_cyberark.py index 87bcd2ffb1c..375b997e02c 100644 --- a/tests/e2e/secret_manager/secret_store_cyberark.py +++ b/tests/e2e/secret_manager/secret_store_cyberark.py @@ -2,6 +2,7 @@ from __future__ import annotations import base64 import os +import time from dataclasses import dataclass, field from typing import Final, Literal from urllib.parse import quote @@ -25,6 +26,9 @@ DEFAULT_USERNAME: Final = "admin" SYSTEM: Final = "cyberark" +_POLICY_LOAD_ATTEMPTS: Final = 5 +_POLICY_LOAD_RETRY_DELAY_SECONDS: Final = 0.2 + _START_HINT: Final = ( f"Start one with `bash tests/e2e/secret_manager/backend.sh up {SYSTEM}`, which writes the env for " f"the proxy (booted from gateway/secret_manager_{SYSTEM}_ci_config.yml) and for the tests" @@ -72,13 +76,20 @@ class Conjur: def _secret_url(self, name: str) -> str: return f"{self.base_url}/secrets/{self.account}/variable/{quote(name, safe='')}" - def _update_root_policy(self, method: Literal["POST", "PATCH"], policy: str, action: str) -> None: + def _load_root_policy(self, method: Literal["POST", "PATCH"], policy: str, attempt: int = 0) -> ExternalWrite: result: Final = send_text_external( method, f"{self.base_url}/policies/{self.account}/policy/root", headers=self._headers(content_type="application/x-yaml"), content=policy, ) + if result.status_code != 409 or attempt + 1 == _POLICY_LOAD_ATTEMPTS: + return result + time.sleep(_POLICY_LOAD_RETRY_DELAY_SECONDS * (1 << attempt)) + return self._load_root_policy(method, policy, attempt + 1) + + def _update_root_policy(self, method: Literal["POST", "PATCH"], policy: str, action: str) -> None: + result: Final = self._load_root_policy(method, policy) self._fail_unless_reached(result, action) if not result.ok: pytest.fail(f"Conjur refused to {action}: HTTP {result.status_code} {result.body[:300]}") diff --git a/tests/e2e/ui/tests/auth/logout.spec.ts b/tests/e2e/ui/tests/auth/logout.spec.ts index 92c31456353..fcdd71898e2 100644 --- a/tests/e2e/ui/tests/auth/logout.spec.ts +++ b/tests/e2e/ui/tests/auth/logout.spec.ts @@ -24,6 +24,11 @@ test.describe("Logout", () => { // Click Logout — the handler clears the auth cookie and navigates via // window.location.href = PROXY_LOGOUT_URL (empty string in the e2e env). await popup.getByRole("button", { name: "Logout" }).click(); + await expect + .poll(async () => (await page.context().cookies()).filter((c) => c.name === "token").length, { + timeout: 15_000, + }) + .toBe(0); // The cookie is now gone — visiting a protected page must redirect to /ui/login. await page.goto("/ui?page=llm-playground", { waitUntil: "domcontentloaded" }); diff --git a/tests/integration/_support/database.py b/tests/integration/_support/database.py index 461cdbda1ee..3ad4f205e12 100644 --- a/tests/integration/_support/database.py +++ b/tests/integration/_support/database.py @@ -1,7 +1,12 @@ import os +import uuid +from collections.abc import Generator +from contextlib import contextmanager from typing import Final, LiteralString +from urllib.parse import urlsplit, urlunsplit import psycopg +from psycopg import sql from psycopg.rows import dict_row from pydantic import JsonValue, TypeAdapter @@ -19,3 +24,14 @@ def read_rows( def write_rows(query: LiteralString, parameters: tuple[str, ...], *, database_url: str | None = None) -> None: with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection: connection.execute(query, parameters) + + +@contextmanager +def scratch_database() -> Generator[str]: + name: Final = f"integration_{uuid.uuid4().hex}" + with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as admin: + admin.execute(sql.SQL("CREATE DATABASE {}").format(sql.Identifier(name))) + try: + yield urlunsplit(urlsplit(os.environ["DATABASE_URL"])._replace(path=f"/{name}")) + finally: + admin.execute(sql.SQL("DROP DATABASE {} WITH (FORCE)").format(sql.Identifier(name))) diff --git a/tests/integration/_support/mail.py b/tests/integration/_support/mail.py new file mode 100644 index 00000000000..3894baeccc3 --- /dev/null +++ b/tests/integration/_support/mail.py @@ -0,0 +1,127 @@ +from __future__ import annotations + +import socketserver +import threading +from collections.abc import Generator +from contextlib import contextmanager +from dataclasses import dataclass +from email import message_from_bytes +from email.message import Message +from queue import SimpleQueue +from typing import Final + + +@dataclass(frozen=True, slots=True) +class Delivery: + sender: str + recipients: tuple[str, ...] + message: Message + + @property + def subject(self) -> str: + return str(self.message["Subject"]) + + @property + def html(self) -> str: + for part in self.message.walk(): + if part.get_content_type() == "text/html": + return part.get_payload(decode=True).decode() + return "" + + +class Mailbox: + def __init__(self, host: str, port: int) -> None: + self.host: Final = host + self.port: Final = port + self._lock: Final = threading.Lock() + self._deliveries: tuple[Delivery, ...] = () + + def record(self, delivery: Delivery) -> None: + with self._lock: + self._deliveries = (*self._deliveries, delivery) + + def deliveries(self) -> tuple[Delivery, ...]: + with self._lock: + return self._deliveries + + +def _address(argument: str) -> str: + return argument.split(":", 1)[1].strip().strip("<>") + + +@contextmanager +def smtp_sink() -> Generator[Mailbox, None, None]: + """Owned plaintext SMTP peer; deliveries traverse the proxy's real smtplib client.""" + errors: Final[SimpleQueue[Exception]] = SimpleQueue() + + class Handler(socketserver.StreamRequestHandler): + timeout = 5 + + def handle(self) -> None: + try: + self._session() + except Exception as error: + errors.put(error) + + def _reply(self, line: str) -> None: + self.wfile.write(f"{line}\r\n".encode()) + self.wfile.flush() + + def _session(self) -> None: + self._reply("220 integration-smtp ready") + # rebind-ok: the SMTP envelope is built across MAIL/RCPT lines and reset after DATA or RSET. + sender = "" + recipients: tuple[str, ...] = () + while True: + raw: Final = self.rfile.readline() + if not raw: + return + line: Final = raw.decode().rstrip("\r\n") + verb: Final = line.split(" ", 1)[0].upper() + if verb in {"EHLO", "HELO"}: + self._reply("250 integration-smtp") + elif verb == "MAIL": + sender = _address(line) + self._reply("250 OK") + elif verb == "RCPT": + recipients = (*recipients, _address(line)) + self._reply("250 OK") + elif verb == "DATA": + self._reply("354 End data with .") + body = bytearray() + while True: + chunk: Final = self.rfile.readline() + if not chunk or chunk == b".\r\n": + break + body.extend(chunk[1:] if chunk.startswith(b"..") else chunk) + mailbox.record(Delivery(sender, recipients, message_from_bytes(bytes(body)))) + sender, recipients = "", () + self._reply("250 OK queued") + elif verb == "RSET": + sender, recipients = "", () + self._reply("250 OK") + elif verb == "NOOP": + self._reply("250 OK") + elif verb == "QUIT": + self._reply("221 Bye") + return + else: + self._reply("502 Command not implemented") + + class OwnedServer(socketserver.ThreadingTCPServer): + allow_reuse_address = True + daemon_threads = False + + with OwnedServer(("127.0.0.1", 0), Handler) as server: + mailbox: Final = Mailbox("127.0.0.1", server.server_address[1]) + thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05}) + thread.start() + try: + yield mailbox + finally: + server.shutdown() + thread.join(timeout=6) + assert not thread.is_alive(), "Owned SMTP server survived cleanup" + server.server_close() + failure: Final = None if errors.empty() else errors.get_nowait() + assert failure is None, f"Owned SMTP peer failed: {failure!r}" diff --git a/tests/integration/mcp/test_mcp_llm_endpoints.py b/tests/integration/mcp/test_mcp_llm_endpoints.py index 6beea9ae8f4..26f5ced4de7 100644 --- a/tests/integration/mcp/test_mcp_llm_endpoints.py +++ b/tests/integration/mcp/test_mcp_llm_endpoints.py @@ -47,6 +47,8 @@ def _model_double(tool: str) -> Callable[[Request], Reply]: arguments: Final = json.dumps(ADD) def respond(request: Request) -> Reply: + if request.method == "GET" and request.target.endswith("/models"): + return _json({"object": "list", "data": []}) body: Final = json.loads(request.body) assert isinstance(body, dict), request.body done: Final = _has_tool_result(body) @@ -192,7 +194,9 @@ class Rig: ) def upstream_tools(self) -> tuple[tuple[str, ...], ...]: - return tuple(_tool_names(json.loads(request.body)) for request in self.wire.drain()) + return tuple( + _tool_names(json.loads(request.body)) for request in self.wire.drain() if request.method == "POST" + ) def final_text(self, body: Mapping[str, object]) -> str: if self.surface == "chat": diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py index 5a3ffdb9965..13be5a95887 100644 --- a/tests/integration/observability/test_langfuse_delivery.py +++ b/tests/integration/observability/test_langfuse_delivery.py @@ -7,19 +7,24 @@ from pathlib import Path from typing import Final import yaml -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.database import read_rows, scratch_database from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, Wire, wire_server from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest from opentelemetry.proto.common.v1.common_pb2 import KeyValue from opentelemetry.proto.trace.v1.trace_pb2 import Span -from pydantic import BaseModel, TypeAdapter +from pydantic import BaseModel, JsonValue, TypeAdapter PUBLIC_KEY: Final = "pk-lf-integration" SECRET_KEY: Final = "sk-lf-integration" PROJECTS_PATH: Final = "/api/public/projects" TRACES_PATH: Final = "/api/public/otel/v1/traces" PROMPTS_PATH: Final = "/api/public/v2/prompts/" +STOCK_CONFIG: Final = Path("tests/integration/proxy_config.yaml") +CONFIG_SECTIONS: Final = ("litellm_settings", "environment_variables") +LANGFUSE_ENVIRONMENT: Final = ("LANGFUSE_HOST", "LANGFUSE_PUBLIC_KEY", "LANGFUSE_SECRET_KEY") +INHERITED_ENVIRONMENT: Final = (*LANGFUSE_ENVIRONMENT, "DATABASE_URL_READ_REPLICA") _PROXY_CONFIG: Final = TypeAdapter(dict[str, object]) _SETTINGS: Final = TypeAdapter(dict[str, object]) @@ -64,9 +69,7 @@ def _text_prompt(name: str) -> Reply: def _langfuse_config(tmp_path: Path) -> Path: - config: Final = _PROXY_CONFIG.validate_python( - yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) - ) + config: Final = _PROXY_CONFIG.validate_python(yaml.safe_load(STOCK_CONFIG.read_text())) settings: Final = { **_SETTINGS.validate_python(config["litellm_settings"]), "success_callback": ["langfuse"], @@ -86,6 +89,14 @@ def _langfuse_environment(langfuse: Wire) -> dict[str, str]: } +def _config_rows(database_url: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT param_name, param_value FROM "LiteLLM_Config" WHERE param_name IN (%s, %s) ORDER BY param_name', + CONFIG_SECTIONS, + database_url=database_url, + ) + + def _attribute(entries: Sequence[KeyValue], key: str) -> str | list[str] | None: for entry in entries: if entry.key != key: @@ -192,6 +203,94 @@ def test_langfuse_callback_delivers_the_generation_over_otlp_v4_with_the_caller_ ) +def test_langfuse_callback_stored_in_the_db_through_config_update_delivers_the_generation_over_otlp_v4( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = "langfusedb" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + public_key: Final = "pk-lf-db-" + marker + secret_key: Final = "sk-lf-db-" + marker + stock_settings: Final = _SETTINGS.validate_python( + _PROXY_CONFIG.validate_python(yaml.safe_load(STOCK_CONFIG.read_text()))["litellm_settings"] + ) + assert "langfuse" not in json.dumps( + [stock_settings.get(key) for key in ("callbacks", "success_callback", "failure_callback")] + ) + + def upstream(request: Request) -> Reply: + assert request.headers["authorization"] == f"Bearer {provider_secret}" + return _completion(marker + "-answer") + + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + scratch_database() as scratch_url, + wire_server(upstream) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, + tmp_path, + {"DATABASE_URL": scratch_url, "LANGFUSE_FLUSH_INTERVAL": "1"}, + remove_environment=INHERITED_ENVIRONMENT, + ) as candidate, + candidate.scenario() as scenario, + ): + candidate.post( + "/config/update", + { + "litellm_settings": {"success_callback": ["langfuse"]}, + "environment_variables": { + "LANGFUSE_HOST": destination.url, + "LANGFUSE_PUBLIC_KEY": public_key, + "LANGFUSE_SECRET_KEY": secret_key, + }, + }, + ) + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + body: Final = candidate.post( + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": marker + "-question"}], + "metadata": {"generation_name": marker}, + "cache": {"no-cache": True}, + }, + ) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + + def exported() -> tuple[Span, ...]: + received.extend(destination.drain()) + return tuple(span for span in _spans(received) if span.name == marker) + + spans: Final = eventually(exported, lambda values: len(values) == 1, seconds=20) + posts: Final = tuple(request for request in received if request.method == "POST") + assert {request.target for request in posts} == {TRACES_PATH}, [request.target for request in received] + basic: Final = "Basic " + base64.b64encode(f"{public_key}:{secret_key}".encode()).decode() + for request in posts: + assert request.headers["authorization"] == basic + assert request.headers["content-type"] == "application/x-protobuf" + assert request.headers["x-langfuse-ingestion-version"] == "4" + assert provider_secret.encode() not in request.body + assert candidate.key.encode() not in request.body + + attributes: Final = spans[0].attributes + assert _attribute(attributes, "langfuse.observation.type") == "generation" + assert _attribute(attributes, "langfuse.observation.metadata.response_id") == string_value(body["id"]) + assert marker + "-question" in str(_attribute(attributes, "langfuse.observation.input")) + assert marker + "-answer" in str(_attribute(attributes, "langfuse.observation.output")) + + stored: Final = {string_value(row["param_name"]): row["param_value"] for row in _config_rows(scratch_url)} + callbacks: Final = TypeAdapter(list[str]).validate_python( + object_value(stored["litellm_settings"]).get("success_callback") or [] + ) + assert "langfuse" in callbacks, stored + assert set(object_value(stored["environment_variables"])) >= set(LANGFUSE_ENVIRONMENT), stored + assert secret_key not in json.dumps(stored["environment_variables"]), stored + + def test_prompt_fetch_encodes_the_name_retries_a_5xx_once_and_keeps_langfuse_headers_off_the_client( gateway: Gateway, tmp_path: Path ) -> None: diff --git a/tests/integration/observability/test_passthrough_upstream_error_chaos.py b/tests/integration/observability/test_passthrough_upstream_error_chaos.py index e424e8597dc..74770f827dd 100644 --- a/tests/integration/observability/test_passthrough_upstream_error_chaos.py +++ b/tests/integration/observability/test_passthrough_upstream_error_chaos.py @@ -5,6 +5,7 @@ import signal import threading import time import uuid +from dataclasses import dataclass from pathlib import Path from typing import Final @@ -54,6 +55,16 @@ def _error_information(call_id: str) -> dict[str, JsonValue]: return object_value(parsed["error_information"]) +@dataclass(frozen=True, slots=True) +class _Served: + response: httpx.Response + client_port: int + + +def _spend_rows(call_id: str) -> list[dict[str, JsonValue]]: + return read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)) + + def _single_spend_row(call_id: str) -> None: rows: Final = eventually( lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)), @@ -65,19 +76,23 @@ def _single_spend_row(call_id: str) -> None: async def _fire_burst( base_url: str, key: str, count: int, *, tolerate_transport_errors: bool = False -) -> tuple[httpx.Response, ...]: - async def one(client: httpx.AsyncClient, index: int) -> httpx.Response: +) -> tuple[_Served, ...]: + async def one(client: httpx.AsyncClient, index: int) -> _Served: if index % 3 == 0: path: Final = "/gemini/v1beta/models/nope-9:generateContent" elif index % 3 == 1: path = "/gemini/v1beta/models/nope-9:streamGenerateContent?alt=sse" else: path = "/gemini/v1beta/models/healthy-model:streamGenerateContent?alt=sse" - return await client.post( + async with client.stream( + "POST", path, json=_GENERATE_CONTENT, headers={"Authorization": f"Bearer {key}", "x-goog-api-key": key}, - ) + ) as response: + client_port: Final = int(response.extensions["network_stream"].get_extra_info("client_addr")[1]) + await response.aread() + return _Served(response=response, client_port=client_port) async with httpx.AsyncClient(base_url=base_url, timeout=30, trust_env=False) as client: results: Final = await asyncio.gather( @@ -85,7 +100,7 @@ async def _fire_burst( ) for result in results: assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) - return tuple(result for result in results if isinstance(result, httpx.Response)) + return tuple(result for result in results if isinstance(result, _Served)) async def test_passthrough_upstream_outage_mid_burst_still_logs_errors_once(gateway: Gateway, tmp_path: Path) -> None: @@ -100,7 +115,7 @@ async def test_passthrough_upstream_outage_mid_burst_still_logs_errors_once(gate burst: Final = asyncio.create_task(_fire_burst(str(candidate.client.base_url), candidate.key, 30)) await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 10, 30) with wire_server(_chaos_reply, port=port): - responses: Final = await burst + responses: Final = tuple(served.response for served in await burst) assert len(responses) == 30 for response in responses: assert response.status_code in (200, 404, 500, 502), response.status_code @@ -137,10 +152,15 @@ async def test_passthrough_worker_sigkill_leaves_sibling_serving_and_logging(gat _fire_burst(str(candidate.client.base_url), candidate.key, 20, tolerate_transport_errors=True) ) await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 5, 30) - psutil.Process(workers[0]).send_signal(signal.SIGKILL) - responses: Final = await burst - for response in responses: - assert response.status_code in (200, 404, 500, 502), response.status_code + victim: Final = psutil.Process(workers[0]) + victim.suspend() + victim_ports: Final = frozenset( + connection.raddr.port for connection in victim.net_connections(kind="tcp") if connection.raddr + ) + victim.send_signal(signal.SIGKILL) + served: Final = await burst + for item in served: + assert item.response.status_code in (200, 404, 500, 502), item.response.status_code follow_up: Final = candidate.request( "POST", "/gemini/v1beta/models/nope-9:generateContent", @@ -149,9 +169,13 @@ async def test_passthrough_worker_sigkill_leaves_sibling_serving_and_logging(gat ) assert follow_up.status_code == 404, follow_up.text assert follow_up.json() == json.loads(_NOT_FOUND_BODY), follow_up.text - for response in responses: - if "x-litellm-call-id" in response.headers: - _single_spend_row(response.headers["x-litellm-call-id"]) + logged: Final = tuple(item for item in served if "x-litellm-call-id" in item.response.headers) + survivor_served: Final = tuple(item for item in logged if item.client_port not in victim_ports) + assert survivor_served, [item.client_port for item in logged] + for item in survivor_served: + _single_spend_row(item.response.headers["x-litellm-call-id"]) + for item in logged: + assert len(_spend_rows(item.response.headers["x-litellm-call-id"])) <= 1, item.response.headers error_information: Final = _error_information(follow_up.headers["x-litellm-call-id"]) assert "not found for this scripted upstream" in str(error_information["error_message"]), follow_up.text diff --git a/tests/integration/sandbox/test_e2b_sandbox.py b/tests/integration/sandbox/test_e2b_sandbox.py index d1cff2ce178..5adfb99db61 100644 --- a/tests/integration/sandbox/test_e2b_sandbox.py +++ b/tests/integration/sandbox/test_e2b_sandbox.py @@ -2,8 +2,7 @@ e2b code execution sandbox - end-to-end integration tests. These tests make REAL HTTP calls to the e2b API and are skipped automatically -unless E2B_API_KEY is set. Mock-only unit tests live in -tests/test_litellm/sandbox/test_e2b_sandbox.py. +unless E2B_API_KEY is set. Run only these tests: pytest tests/integration/sandbox/test_e2b_sandbox.py -v diff --git a/tests/integration/spend/test_cache_and_quota.py b/tests/integration/spend/test_cache_and_quota.py index 1c2b5551855..1cc03838f8d 100644 --- a/tests/integration/spend/test_cache_and_quota.py +++ b/tests/integration/spend/test_cache_and_quota.py @@ -1,27 +1,22 @@ import json -import os import threading import uuid -from collections.abc import Generator from concurrent.futures import ThreadPoolExecutor -from contextlib import ExitStack, contextmanager +from contextlib import ExitStack from hashlib import sha256 from pathlib import Path from typing import Final -from urllib.parse import urlsplit, urlunsplit import httpx -import psycopg import pytest from hypothesis import strategies as st from hypothesis.stateful import RuleBasedStateMachine, rule, run_state_machine_as_test from integration._support.client import Gateway, eventually, string_value -from integration._support.database import read_rows +from integration._support.database import read_rows, scratch_database from integration._support.database_relay import database_relay from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, wire_server -from psycopg import sql @pytest.mark.covers("quota_management.response_cache.generated_sequences_preserve_content_and_accounting") @@ -225,17 +220,6 @@ def test_key_budget_at_boundary_blocks_provider_then_explicit_reset_restores(gat RESET_SWEEP_QUERY: Final = b'"LiteLLM_VerificationToken"."budget_reset_at" < $' -@contextmanager -def scratch_database() -> Generator[str]: - name: Final = f"integration_{uuid.uuid4().hex}" - with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as admin: - admin.execute(sql.SQL("CREATE DATABASE {}").format(sql.Identifier(name))) - try: - yield urlunsplit(urlsplit(os.environ["DATABASE_URL"])._replace(path=f"/{name}")) - finally: - admin.execute(sql.SQL("DROP DATABASE {} WITH (FORCE)").format(sql.Identifier(name))) - - @pytest.mark.covers("quota_management.budget.key.scheduled_reset_survives_transient_db_outage") @pytest.mark.timeout(300) def test_scheduled_budget_reset_reconnects_after_db_transport_failure_and_unblocks_key( diff --git a/tests/integration/spend/test_team_member_budget_alerts.py b/tests/integration/spend/test_team_member_budget_alerts.py new file mode 100644 index 00000000000..f12bcb9748a --- /dev/null +++ b/tests/integration/spend/test_team_member_budget_alerts.py @@ -0,0 +1,93 @@ +import uuid +from pathlib import Path +from typing import Final + +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.mail import smtp_sink +from integration._support.process import owned_proxy + +MEMBER_BUDGET: Final = 0.10 +CALL_COST: Final = 20 * 0.001 + 20 * 0.002 + + +def _membership_spend(user_id: str, team_id: str) -> float: + rows: Final = read_rows( + 'SELECT spend FROM "LiteLLM_TeamMembership" WHERE user_id = %s AND team_id = %s', (user_id, team_id) + ) + return float(str(rows[0]["spend"])) if rows else 0.0 + + +def test_team_member_budget_thresholds_email_member_and_configured_recipients(gateway: Gateway, tmp_path: Path) -> None: + member_email: Final = f"member-{uuid.uuid4().hex}@integration.test" + finance_email: Final = f"finance-{uuid.uuid4().hex}@integration.test" + configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + configuration["general_settings"]["alerting"] = ["email"] + path: Final = tmp_path / "email-alerting.yaml" + path.write_text(yaml.safe_dump(configuration)) + with smtp_sink() as mailbox: + overrides: Final = { + "SMTP_HOST": mailbox.host, + "SMTP_PORT": str(mailbox.port), + "SMTP_TLS": "False", + "SMTP_SENDER_EMAIL": "alerts@integration.test", + } + with owned_proxy(gateway, tmp_path, overrides, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002) + user_id: Final = scenario.user(user_email=member_email) + team_id: Final = scenario.team( + models=[model], + team_member_budget=MEMBER_BUDGET, + metadata={"team_member_max_budget_alert_emails": {"50": [], "100": [finance_email]}}, + ) + candidate.post("/team/member_add", {"team_id": team_id, "member": {"user_id": user_id, "role": "user"}}) + key: Final = scenario.key(team_id=team_id, user_id=user_id) + + first: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "first call"}]}, + key=key, + ) + assert first.status_code == 200, first.text + assert float(first.headers["x-litellm-response-cost"]) == pytest.approx(CALL_COST) + eventually( + lambda: _membership_spend(user_id, team_id), lambda spend: spend == pytest.approx(CALL_COST), seconds=70 + ) + assert mailbox.deliveries() == (), "no threshold is reached before the first call is recorded" + + second: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "second call"}]}, + key=key, + ) + assert second.status_code == 200, second.text + halfway: Final = eventually(mailbox.deliveries, lambda found: len(found) >= 1, seconds=30) + assert [delivery.recipients for delivery in halfway] == [(member_email,)], halfway + assert "50%" in halfway[0].subject, halfway[0].subject + assert f"${MEMBER_BUDGET}" in halfway[0].html, halfway[0].html + eventually( + lambda: _membership_spend(user_id, team_id), + lambda spend: spend == pytest.approx(2 * CALL_COST), + seconds=70, + ) + + third: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "third call"}]}, + key=key, + ) + assert third.status_code == 422 and third.json()["error"]["type"] == "budget_exceeded", third.text + capped: Final = eventually(mailbox.deliveries, lambda found: len(found) >= 3, seconds=30) + hundred: Final = capped[1:] + assert all("100%" in delivery.subject for delivery in hundred), capped + assert {recipient for delivery in hundred for recipient in delivery.recipients} == { + member_email, + finance_email, + }, capped + assert all(member_email in delivery.html and f"${MEMBER_BUDGET}" in delivery.html for delivery in hundred) + assert len(capped) == 3, capped diff --git a/tests/litellm_utils_tests/test_cyberark.py b/tests/litellm_utils_tests/test_cyberark.py index 9172e33af10..6d52cd9b079 100644 --- a/tests/litellm_utils_tests/test_cyberark.py +++ b/tests/litellm_utils_tests/test_cyberark.py @@ -86,7 +86,8 @@ async def test_cyberark_write_secret_rejects_yaml_injection(): "team/user@example.com", ], ) -def test_cyberark_ensure_variable_exists_escapes_yaml_metacharacters(secret_name): +@pytest.mark.asyncio +async def test_cyberark_ensure_variable_exists_escapes_yaml_metacharacters(secret_name): """ Regression test: _ensure_variable_exists must escape secret_name (not just denylist-check it) so the policy body always parses back to exactly one @@ -95,19 +96,21 @@ def test_cyberark_ensure_variable_exists_escapes_yaml_metacharacters(secret_name with patch("litellm.proxy.proxy_server.premium_user", True): captured = {} - def _capture_post(url, headers=None, content=None): + async def _capture_post(url, headers=None, content=None): captured["content"] = content return create_mock_response(status_code=201, text="") mock_sync_client = MagicMock() - mock_sync_client.client.post.side_effect = _capture_post + mock_sync_client.client.post.return_value = create_mock_response(status_code=200, text="mock-token") + mock_async_client = MagicMock() + mock_async_client.client.post.side_effect = _capture_post with patch( "litellm.secret_managers.cyberark_secret_manager._get_httpx_client", return_value=mock_sync_client, ): cyberark_manager = CyberArkSecretManager() - cyberark_manager._ensure_variable_exists(secret_name) + await cyberark_manager._ensure_variable_exists(secret_name, mock_async_client) policy_yaml = captured["content"] parsed = yaml.compose(policy_yaml) diff --git a/tests/llm_translation/base_embedding_unit_tests.py b/tests/llm_translation/base_embedding_unit_tests.py index 1a88f0e9d6b..469416fc0cf 100644 --- a/tests/llm_translation/base_embedding_unit_tests.py +++ b/tests/llm_translation/base_embedding_unit_tests.py @@ -16,15 +16,12 @@ from litellm.utils import ( get_optional_params, get_optional_params_embeddings, ) -import requests import base64 +from pathlib import Path -# test_example.py from abc import ABC, abstractmethod -url = "https://dummyimage.com/100/100/fff&text=Test+image" -response = requests.get(url) -file_data = response.content +file_data = (Path(__file__).parent.parent / "white_100x100.png").read_bytes() encoded_file = base64.b64encode(file_data).decode("utf-8") base64_image = f"data:image/png;base64,{encoded_file}" diff --git a/tests/test_litellm/llms/databricks/databricks_config.template.txt b/tests/llm_translation/databricks_config.template.txt similarity index 100% rename from tests/test_litellm/llms/databricks/databricks_config.template.txt rename to tests/llm_translation/databricks_config.template.txt diff --git a/tests/test_litellm/interactions/base_interactions_test.py b/tests/llm_translation/interactions/base_interactions_test.py similarity index 100% rename from tests/test_litellm/interactions/base_interactions_test.py rename to tests/llm_translation/interactions/base_interactions_test.py diff --git a/tests/test_litellm/interactions/test_gemini_interactions.py b/tests/llm_translation/interactions/test_gemini_interactions.py similarity index 88% rename from tests/test_litellm/interactions/test_gemini_interactions.py rename to tests/llm_translation/interactions/test_gemini_interactions.py index afce77e3ce4..0ab4da952e6 100644 --- a/tests/test_litellm/interactions/test_gemini_interactions.py +++ b/tests/llm_translation/interactions/test_gemini_interactions.py @@ -6,7 +6,7 @@ Inherits from BaseInteractionsTest to run the same test suite against Gemini. import os -from tests.test_litellm.interactions.base_interactions_test import ( +from tests.llm_translation.interactions.base_interactions_test import ( BaseInteractionsTest, ) diff --git a/tests/test_litellm/interactions/test_google_interactions_integration.py b/tests/llm_translation/interactions/test_google_interactions_integration.py similarity index 99% rename from tests/test_litellm/interactions/test_google_interactions_integration.py rename to tests/llm_translation/interactions/test_google_interactions_integration.py index 93429d64789..10e6cf86e6d 100644 --- a/tests/test_litellm/interactions/test_google_interactions_integration.py +++ b/tests/llm_translation/interactions/test_google_interactions_integration.py @@ -5,7 +5,7 @@ Tests the litellm.interactions.create() and related methods against the Google A Per OpenAPI spec: https://ai.google.dev/static/api/interactions.openapi.json -Run with: pytest tests/test_litellm/interactions/test_google_interactions_integration.py -v +Run with: pytest tests/llm_translation/interactions/test_google_interactions_integration.py -v """ import asyncio diff --git a/tests/test_litellm/interactions/test_litellm_responses_bridge.py b/tests/llm_translation/interactions/test_litellm_responses_bridge.py similarity index 91% rename from tests/test_litellm/interactions/test_litellm_responses_bridge.py rename to tests/llm_translation/interactions/test_litellm_responses_bridge.py index 17e7f9fc4ff..ae025ab60b0 100644 --- a/tests/test_litellm/interactions/test_litellm_responses_bridge.py +++ b/tests/llm_translation/interactions/test_litellm_responses_bridge.py @@ -7,7 +7,7 @@ the litellm_responses bridge provider, which calls litellm.responses() internall import os -from tests.test_litellm.interactions.base_interactions_test import ( +from tests.llm_translation.interactions.base_interactions_test import ( BaseInteractionsTest, ) diff --git a/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py b/tests/llm_translation/test_cometapi_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py rename to tests/llm_translation/test_cometapi_chat_transformation.py diff --git a/tests/test_litellm/test_compression.py b/tests/llm_translation/test_compression.py similarity index 100% rename from tests/test_litellm/test_compression.py rename to tests/llm_translation/test_compression.py diff --git a/tests/test_litellm/llms/databricks/test_databricks_e2e.py b/tests/llm_translation/test_databricks_e2e.py similarity index 99% rename from tests/test_litellm/llms/databricks/test_databricks_e2e.py rename to tests/llm_translation/test_databricks_e2e.py index 669f9e94639..a979988102e 100644 --- a/tests/test_litellm/llms/databricks/test_databricks_e2e.py +++ b/tests/llm_translation/test_databricks_e2e.py @@ -51,7 +51,7 @@ Setup: Run with: cd /path/to/litellm - python tests/test_litellm/llms/databricks/test_databricks_e2e.py + python tests/llm_translation/test_databricks_e2e.py Config Options: TEST_AUTH_METHOD=oauth # Test OAuth M2M authentication @@ -69,12 +69,12 @@ import pytest # These are E2E tests that require real Databricks credentials pytestmark = pytest.mark.skip( reason="E2E tests require real Databricks credentials. Run directly with: " - "python tests/test_litellm/llms/databricks/test_databricks_e2e.py" + "python tests/llm_translation/test_databricks_e2e.py" ) # Add the litellm package to path sys.path.insert( - 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../..")) + 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")) ) # Config file path - can be overridden with DATABRICKS_TEST_CONFIG env var diff --git a/tests/test_litellm/llms/openai_like/test_json_providers.py b/tests/llm_translation/test_json_providers.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_json_providers.py rename to tests/llm_translation/test_json_providers.py diff --git a/tests/test_litellm/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py b/tests/llm_translation/test_mistral_audio_transcription_transformation.py similarity index 100% rename from tests/test_litellm/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py rename to tests/llm_translation/test_mistral_audio_transcription_transformation.py diff --git a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py b/tests/llm_translation/test_ovhcloud_audio_transcription_transformation.py similarity index 100% rename from tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py rename to tests/llm_translation/test_ovhcloud_audio_transcription_transformation.py diff --git a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py b/tests/llm_translation/test_ovhcloud_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py rename to tests/llm_translation/test_ovhcloud_chat_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py b/tests/llm_translation/test_vertex_ai_image_generation_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py rename to tests/llm_translation/test_vertex_ai_image_generation_transformation.py diff --git a/tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py b/tests/llm_translation/test_xiaomi_mimo.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py rename to tests/llm_translation/test_xiaomi_mimo.py diff --git a/tests/local_testing/conftest.py b/tests/local_testing/conftest.py index d03f074f557..df3dacac3b2 100644 --- a/tests/local_testing/conftest.py +++ b/tests/local_testing/conftest.py @@ -19,6 +19,7 @@ import pytest import litellm from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER +from litellm.utils import _invalidate_model_cost_lowercase_map # ``litellm.model_cost`` is loaded at import time from the URL pinned to ``main`` # (``LITELLM_MODEL_COST_MAP_URL``). The in-tree backup ships with this branch @@ -232,6 +233,7 @@ def isolate_litellm_state(): for attr, original_value in original_state.items(): if hasattr(litellm, attr): setattr(litellm, attr, original_value) + _invalidate_model_cost_lowercase_map() @pytest.fixture(scope="module", autouse=True) diff --git a/tests/local_testing/test_alangfuse.py b/tests/local_testing/test_alangfuse.py index a9d111843fd..ec80724d3ba 100644 --- a/tests/local_testing/test_alangfuse.py +++ b/tests/local_testing/test_alangfuse.py @@ -5,6 +5,10 @@ import logging import os from typing import Any, Optional from unittest.mock import MagicMock, patch +import threading +from http.server import BaseHTTPRequestHandler, HTTPServer + +from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest logging.basicConfig(level=logging.DEBUG) @@ -206,53 +210,91 @@ def create_async_task(**completion_kwargs): return asyncio.create_task(litellm.acompletion(**completion_args)) +def _otlp_capture(exports: list[bytes]) -> type[BaseHTTPRequestHandler]: + class OtlpCapture(BaseHTTPRequestHandler): + def do_POST(self): + exports.append(self.rfile.read(int(self.headers.get("content-length", 0)))) + self.send_response(200) + self.end_headers() + + def do_GET(self): + self.send_response(200) + self.send_header("content-type", "application/json") + self.end_headers() + self.wfile.write(b"{}") + + def log_message(self, *args): + pass + + return OtlpCapture + + +@pytest.fixture +def local_langfuse(): + exports: list[bytes] = [] + server = HTTPServer(("127.0.0.1", 0), _otlp_capture(exports)) + threading.Thread(target=server.serve_forever, daemon=True).start() + yield f"http://127.0.0.1:{server.server_port}", exports + server.shutdown() + + +def _exported_spans(exports: list[bytes]): + for body in exports: + for resource_spans in ExportTraceServiceRequest.FromString(body).resource_spans: + for scope_spans in resource_spans.scope_spans: + yield from scope_spans.spans + + +def _exported_attributes(exports: list[bytes], trace_id: str) -> list[dict[str, str]]: + return [ + {attribute.key: attribute.value.string_value for attribute in span.attributes} + for span in _exported_spans(list(exports)) + if span.trace_id.hex() == trace_id + ] + + @pytest.mark.asyncio @pytest.mark.parametrize("stream", [False, True]) -@pytest.mark.flaky(retries=12, delay=2) -async def test_langfuse_logging_without_request_response(stream, langfuse_client): - try: - from litellm._uuid import uuid +async def test_langfuse_logging_without_request_response(stream, local_langfuse, monkeypatch): + from litellm._uuid import uuid - _unique_trace_name = f"litellm-test-{str(uuid.uuid4())}" - litellm.set_verbose = True - litellm.turn_off_message_logging = True - litellm.success_callback = ["langfuse"] - response = await create_async_task( - model="gpt-3.5-turbo", - stream=stream, - metadata={"trace_id": _unique_trace_name}, - ) - print(response) - if stream: - async for chunk in response: - print(chunk) + langfuse_host, exports = local_langfuse + prompt = f"prompt-{uuid.uuid4()}" + answer = f"answer-{uuid.uuid4()}" + trace_name = f"litellm-test-{uuid.uuid4()}" + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + monkeypatch.setattr(litellm, "success_callback", ["langfuse"]) + response = await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": prompt}], + mock_response=answer, + stream=stream, + metadata={"trace_id": trace_name}, + langfuse_public_key=f"pk-lf-{trace_name}", + langfuse_secret_key="sk-lf-local", + langfuse_host=langfuse_host, + ) + if stream: + async for _ in response: + pass - langfuse_client.flush() + generations: list[dict[str, str]] = [] + for _ in range(60): + generations = [ + attributes + for attributes in _exported_attributes(exports, resolve_trace_id(trace_name)) + if attributes.get("langfuse.observation.type") == "generation" + ] + if generations: + break + await asyncio.sleep(0.5) - for _ in range(30): - _trace_data = langfuse_client.api.observations.get_many( - trace_id=resolve_trace_id(_unique_trace_name), - type="GENERATION", - fields="core,io", - ).data - if _trace_data: - break - await asyncio.sleep(3) - - print(f"_trace_data: {_trace_data}") - assert json.loads(_trace_data[0].input) == { - "messages": [{"content": "redacted-by-litellm", "role": "user"}] - } - assert json.loads(_trace_data[0].output) == { - "role": "assistant", - "content": "redacted-by-litellm", - "function_call": None, - "tool_calls": None, - "provider_specific_fields": None, - } - - except Exception as e: - pytest.fail(f"An exception occurred - {e}") + assert len(generations) == 1, generations + assert json.loads(generations[0]["langfuse.observation.input"]) == { + "messages": [{"content": "redacted-by-litellm", "role": "user"}] + } + assert json.loads(generations[0]["langfuse.observation.output"])["content"] == "redacted-by-litellm" + assert all(prompt.encode() not in body and answer.encode() not in body for body in exports) # Get the current directory of the file being run diff --git a/tests/local_testing/test_get_model_info.py b/tests/local_testing/test_get_model_info.py index 1e46a1bf853..79f6739a423 100644 --- a/tests/local_testing/test_get_model_info.py +++ b/tests/local_testing/test_get_model_info.py @@ -9,6 +9,7 @@ import pytest import litellm from litellm import get_model_info +from litellm.utils import _invalidate_model_cost_lowercase_map from unittest.mock import MagicMock, patch @@ -74,15 +75,15 @@ def test_get_model_info_ollama_chat(): assert mock_client.call_args.kwargs["json"]["name"] == "unknown-model" -def test_get_model_info_bedrock_region(): - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - args = { - "model": "us.anthropic.claude-haiku-4-5-20251001-v1:0", - "custom_llm_provider": "bedrock", +def test_get_model_info_bedrock_region(monkeypatch): + regional_model = "us.anthropic.claude-haiku-4-5-20251001-v1:0" + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + model_cost_without_regional_entry = { + key: value for key, value in litellm.get_model_cost_map(url="").items() if key != regional_model } - litellm.model_cost.pop("us.anthropic.claude-haiku-4-5-20251001-v1:0", None) - info = litellm.get_model_info(**args) + monkeypatch.setattr(litellm, "model_cost", model_cost_without_regional_entry) + _invalidate_model_cost_lowercase_map() + info = litellm.get_model_info(model=regional_model, custom_llm_provider="bedrock") print("info", info) assert info["key"] == "anthropic.claude-haiku-4-5-20251001-v1:0" assert info["litellm_provider"] == "bedrock_converse" @@ -319,6 +320,33 @@ def test_get_model_info_bedrock_cross_region_capability_parity(): assert checked > 0, "no cross-region bedrock profiles found - the filter is inert" + +def test_get_model_info_bedrock_priced_cross_region_profile_has_priced_base(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + prefixes = ("us.", "eu.", "apac.", "us-gov.", "au.", "global.") + checked = 0 + + for k, v in litellm.model_cost.items(): + if not str(v.get("litellm_provider", "")).startswith("bedrock"): + continue + base_model_key = next( + (k[len(p) :] for p in prefixes if k.startswith(p)), + None, + ) + if base_model_key is None or base_model_key not in litellm.model_cost: + continue + checked += 1 + base = litellm.model_cost[base_model_key] + for cost_key in ("input_cost_per_token", "output_cost_per_token"): + if (v.get(cost_key) or 0) > 0: + assert ( + base.get(cost_key) or 0 + ) > 0, f"{k} charges {cost_key} but its base {base_model_key} is free" + + assert checked > 0, "no cross-region bedrock profiles found - the filter is inert" + def test_get_model_info_huggingface_models(monkeypatch): from litellm import Router from litellm.types.router import ModelGroupInfo diff --git a/tests/local_testing/test_handler_gc_does_not_close_client.py b/tests/local_testing/test_handler_gc_does_not_close_client.py index 63c5694dd89..d6987107fa8 100644 --- a/tests/local_testing/test_handler_gc_does_not_close_client.py +++ b/tests/local_testing/test_handler_gc_does_not_close_client.py @@ -23,10 +23,7 @@ test here may keep the client in a local: that inflates the very refcount under test, and the test then passes on a broken handler. They hold weak references instead, which the refcount does not count. -These live here rather than under ``tests/test_litellm/`` because they need a -real connection pool: a mocked transport goes on yielding chunks after its -client is closed, so the very teardown under test is what a mock cannot -reproduce. The server is a hermetic, credential-free ``ThreadingHTTPServer`` on +The server is a hermetic, credential-free ``ThreadingHTTPServer`` on an ephemeral loopback port, and needs no network access beyond it. Related: https://github.com/BerriAI/litellm/issues/24929 diff --git a/tests/proxy_behavior/management/conftest.py b/tests/proxy_behavior/management/conftest.py index 255b937bdd3..74db323e7f2 100644 --- a/tests/proxy_behavior/management/conftest.py +++ b/tests/proxy_behavior/management/conftest.py @@ -31,7 +31,7 @@ def _write_minimal_proxy_config() -> str: return f.name -@pytest_asyncio.fixture(scope="session") +@pytest_asyncio.fixture(scope="package") async def proxy_app(): from litellm.proxy import proxy_server from litellm.proxy.proxy_server import ( @@ -67,7 +67,7 @@ async def proxy_app(): yield app -@pytest_asyncio.fixture(scope="session") +@pytest_asyncio.fixture(scope="package") async def proxy_client(proxy_app) -> AsyncIterator[httpx.AsyncClient]: transport = httpx.ASGITransport(app=proxy_app) async with httpx.AsyncClient( @@ -76,7 +76,7 @@ async def proxy_client(proxy_app) -> AsyncIterator[httpx.AsyncClient]: yield client -@pytest_asyncio.fixture(scope="session") +@pytest_asyncio.fixture(scope="package") async def prisma(proxy_app): from litellm.proxy import proxy_server @@ -84,7 +84,7 @@ async def prisma(proxy_app): return proxy_server.prisma_client -@pytest_asyncio.fixture(scope="session") +@pytest_asyncio.fixture(scope="package") async def world(prisma): from .actors import seed_world diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index 7e593767ea9..d3ad1d989c8 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -4,7 +4,7 @@ import os import traceback from dotenv import load_dotenv from fastapi import Request -from datetime import datetime +from datetime import datetime, timezone from litellm import Router import pytest @@ -971,11 +971,18 @@ def _rpm_tpm_router(model_id: str) -> Router: ) +@pytest.fixture +def router_minute_pinned(monkeypatch): + pinned = datetime(2026, 1, 1, 12, 0, 30, tzinfo=timezone.utc) + monkeypatch.setattr("litellm.router.get_utc_datetime", lambda: pinned) + + def _ratelimit_headers(response: ModelResponse | CustomStreamWrapper) -> dict[str, int]: return {k: v for k, v in response._hidden_params["additional_headers"].items() if k.startswith("x-ratelimit-")} @pytest.mark.asyncio +@pytest.mark.usefixtures("router_minute_pinned") async def test_acompletion_headers_read_post_increment_counter_and_count_once(): router = _rpm_tpm_router("lit-3058-async") @@ -1018,6 +1025,7 @@ async def test_acompletion_wildcard_route_headers_and_counter_use_resolved_deplo @pytest.mark.asyncio +@pytest.mark.usefixtures("router_minute_pinned") async def test_acompletion_stream_counts_request_before_headers_and_tokens_once_on_completion(): router = _rpm_tpm_router("lit-3058-stream") diff --git a/tests/search_tests/test_bing_grounding_search.py b/tests/search_tests/test_bing_grounding_search.py index f532158e462..6ea79076370 100644 --- a/tests/search_tests/test_bing_grounding_search.py +++ b/tests/search_tests/test_bing_grounding_search.py @@ -196,4 +196,5 @@ class TestBingGroundingSearchTransformation: ): response = litellm.search(query="pricing check", search_provider="bing_grounding") - assert response._hidden_params["response_cost"] == pytest.approx(0.035) + # Grounding with Bing Search (G1 SKU): $14 per 1,000 transactions, https://www.microsoft.com/en-us/bing/apis, checked 2026-09-24 + assert response._hidden_params["response_cost"] == pytest.approx(0.014) diff --git a/tests/store_model_in_db_tests/test_callbacks_in_db.py b/tests/store_model_in_db_tests/test_callbacks_in_db.py deleted file mode 100644 index 6497e4064b7..00000000000 --- a/tests/store_model_in_db_tests/test_callbacks_in_db.py +++ /dev/null @@ -1,114 +0,0 @@ -""" -PROD TEST - DO NOT Delete this Test - -e2e test for langfuse callback in DB -- Add langfuse callback to DB - with /config/update -- wait 20 seconds for the callback to be loaded into the instance -- Make a /chat/completions request to the proxy -- Check if the request is logged in Langfuse -""" - -import pytest -import asyncio -import aiohttp -import os -import dotenv -from dotenv import load_dotenv -from openai import AsyncOpenAI, APIConnectionError -from openai.types.chat import ChatCompletion - -load_dotenv() - -# used for testing -LANGFUSE_BASE_URL = "https://exampleopenaiendpoint-production-c715.up.railway.app" -PROXY_BASE_URL = "http://127.0.0.1:4000" - - -async def wait_for_proxy_ready(session, timeout: int = 60): - for _ in range(timeout): - try: - async with session.get(f"{PROXY_BASE_URL}/health/liveliness") as response: - if response.status == 200: - return - except aiohttp.ClientError: - pass - await asyncio.sleep(1) - raise RuntimeError(f"Proxy at {PROXY_BASE_URL} not ready after {timeout}s") - - -async def config_update(session, routing_strategy=None): - url = f"{PROXY_BASE_URL}/config/update" - headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} - print("routing_strategy: ", routing_strategy) - data = { - "litellm_settings": {"success_callback": ["langfuse"]}, - "environment_variables": { - "LANGFUSE_PUBLIC_KEY": "any-public-key", - "LANGFUSE_SECRET_KEY": "any-secret-key", - "LANGFUSE_HOST": LANGFUSE_BASE_URL, - }, - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print("status: ", status) - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - -async def check_langfuse_request(response_id: str): - async with aiohttp.ClientSession() as session: - url = f"{LANGFUSE_BASE_URL}/langfuse/trace/{response_id}" - async with session.get(url) as response: - response_json = await response.json() - assert response.status == 200, f"Expected status 200, got {response.status}" - assert ( - response_json["exists"] == True - ), f"Request {response_id} not found in Langfuse traces" - assert response_json["request_id"] == response_id, f"Request ID mismatch" - - -async def make_chat_completions_request() -> ChatCompletion: - client = AsyncOpenAI(api_key="sk-1234", base_url=PROXY_BASE_URL) - last_error = None - for _ in range(10): - try: - response = await client.chat.completions.create( - model="fake-openai-endpoint", - messages=[{"role": "user", "content": "Hello, world!"}], - ) - print(response) - return response - except APIConnectionError as e: - last_error = e - await asyncio.sleep(2) - raise AssertionError( - f"Proxy at {PROXY_BASE_URL} unreachable after retries: {last_error!r}" - ) - - -@pytest.mark.asyncio -async def test_e2e_langfuse_callbacks_in_db(): - - async with aiohttp.ClientSession() as session: - # add langfuse callback to DB - await config_update(session) - - # wait 20 seconds for the callback to be loaded into the instance - await asyncio.sleep(20) - await wait_for_proxy_ready(session) - - # make a /chat/completions request to the proxy - response = await make_chat_completions_request() - print(response) - response_id = response.id - print("response_id: ", response_id) - - await asyncio.sleep(20) - # check if the request is logged in Langfuse - await check_langfuse_request(response_id) diff --git a/tests/test_litellm/litellm_core_utils/__init__.py b/tests/test_litellm/litellm_core_utils/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter.py b/tests/test_litellm/litellm_core_utils/test_token_counter.py deleted file mode 100644 index 1e10b7e82b1..00000000000 --- a/tests/test_litellm/litellm_core_utils/test_token_counter.py +++ /dev/null @@ -1,51 +0,0 @@ -import pytest -from litellm import create_pretrained_tokenizer -from tests.unit.litellm_core_utils.test_token_counter import token_counter - - -def test_tokenizers(): - try: - ### test the openai, claude, cohere and llama2 tokenizers. - ### The tokenizer value should be different for all - sample_text = "Hellö World, this is my input string! My name is ishaan CTO" - - # openai tokenizer - openai_tokens = token_counter(model="gpt-3.5-turbo", text=sample_text) - - # claude tokenizer - claude_tokens = token_counter(model="claude-3-5-haiku-20241022", text=sample_text) - - # cohere tokenizer - cohere_tokens = token_counter(model="command-nightly", text=sample_text) - - # llama2 tokenizer - llama2_tokens = token_counter(model="meta-llama/Llama-2-7b-chat", text=sample_text) - - # llama3 tokenizer (also testing custom tokenizer) - llama3_tokens_1 = token_counter(model="meta-llama/llama-3-70b-instruct", text=sample_text) - - try: - llama3_tokenizer = create_pretrained_tokenizer("Xenova/llama-3-tokenizer") - except Exception as e: - pytest.skip(f"custom tokenizer download failed (HF hub unreachable): {e}") - llama3_tokens_2 = token_counter(custom_tokenizer=llama3_tokenizer, text=sample_text) - - print( - f"openai tokens: {openai_tokens}; claude tokens: {claude_tokens}; cohere tokens: {cohere_tokens}; llama2 tokens: {llama2_tokens}; llama3 tokens: {llama3_tokens_1}" - ) - - # assert that all token values are different - # llama2 may fall back to the tiktoken tokenizer when the HuggingFace - # model hub is unreachable (e.g. in CI). In that case the count will - # equal the openai count and the differentiation assertion is skipped. - if openai_tokens == llama2_tokens: - pytest.skip("llama2 fell back to tiktoken (HF hub unreachable); skipping differentiation assertion") - assert llama2_tokens != llama3_tokens_1, "Token values are not different." - - assert llama3_tokens_1 == llama3_tokens_2, ( - "Custom tokenizer is not being used! It has been configured to use the same tokenizer as the built in llama3 tokenizer and the results should be the same." - ) - - print("test tokenizer: It worked!") - except Exception as e: - pytest.fail(f"An exception occured: {e}") diff --git a/tests/test_litellm/litellm_core_utils/test_tokenizer.py b/tests/test_litellm/litellm_core_utils/test_tokenizer.py deleted file mode 100644 index 2171044970c..00000000000 --- a/tests/test_litellm/litellm_core_utils/test_tokenizer.py +++ /dev/null @@ -1,20 +0,0 @@ -import pytest - -from tests.unit.litellm_core_utils.test_tokenizer import ( - UNICODE_TEXTS, - assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface, - assert_openai_encoding_matches_python, -) - -NETWORK_ENCODINGS = ("r50k_base", "gpt2") - - -@pytest.mark.parametrize("name", NETWORK_ENCODINGS) -@pytest.mark.parametrize("text", UNICODE_TEXTS) -def test_openai_encoding_matches_python_unicode_and_batches(name: str, text: str) -> None: - assert_openai_encoding_matches_python(name, text) - - -@pytest.mark.parametrize("name", ("gpt2",)) -def test_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None: - assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name) diff --git a/tests/test_litellm/llms/mistral/__init__.py b/tests/test_litellm/llms/mistral/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/llms/openai_like/__init__.py b/tests/test_litellm/llms/openai_like/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/llms/vertex_ai/__init__.py b/tests/test_litellm/llms/vertex_ai/__init__.py deleted file mode 100644 index fc7e977484b..00000000000 --- a/tests/test_litellm/llms/vertex_ai/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Vertex AI tests package.""" diff --git a/tests/test_litellm/llms/vertex_ai/gemini/__init__.py b/tests/test_litellm/llms/vertex_ai/gemini/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py deleted file mode 100644 index d3a7ba7a1bd..00000000000 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ /dev/null @@ -1,54 +0,0 @@ -import pytest - -from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_result, -) -from litellm.types.llms.vertex_ai import BlobType - - -def test_convert_tool_response_with_url_image(): - """Test tool response with HTTP URL image (will download and convert).""" - # Use a publicly accessible test image URL - test_image_url = "https://via.placeholder.com/1x1.png" - - tool_message = { - "role": "tool", - "tool_call_id": "call_test456", - "content": [ - {"type": "text", "text": '{"url": "https://example.com"}'}, - {"type": "input_image", "image_url": test_image_url}, - ], - } - - last_message_with_tool_calls = { - "tool_calls": [ - { - "id": "call_test456", - "function": { - "name": "type_text_at", - "arguments": '{"x": 300, "y": 400, "text": "hello"}', - }, - } - ] - } - - try: - result = convert_to_gemini_tool_call_result(tool_message, last_message_with_tool_calls) - - assert isinstance(result, list), "Should return a parts list when media is present" - assert len(result) == 1, "Should return one function_response part" - result_part = result[0] - assert "function_response" in result_part - assert "inline_data" not in result_part - function_response = result_part["function_response"] - assert function_response["name"] == "type_text_at" - - # Check inline_data is nested under functionResponse.parts. - assert "parts" in function_response - assert len(function_response["parts"]) == 1 - inline_data: BlobType = function_response["parts"][0]["inline_data"] - assert "data" in inline_data - assert "mime_type" in inline_data - except Exception as e: - # Skip test if URL download fails (no internet connection, etc.) - pytest.skip(f"Failed to download image from URL: {e}") diff --git a/tests/test_litellm/llms/volcengine/__init__.py b/tests/test_litellm/llms/volcengine/__init__.py deleted file mode 100644 index 825e259b1fc..00000000000 --- a/tests/test_litellm/llms/volcengine/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Volcengine tests diff --git a/tests/test_litellm/log.txt b/tests/test_litellm/log.txt deleted file mode 100644 index 6470b12fedb..00000000000 --- a/tests/test_litellm/log.txt +++ /dev/null @@ -1,2 +0,0 @@ -llms/bedrock/chat/invoke_agent/transformation.py:404: error: Incompatible types in assignment (expression has type "object", variable has type "InvokeAgentModelInvocationOutput | None") [assignment] -llms/bedrock/chat/invoke_agent/transformation.py:405: error: Argument 1 to "get" of "Mapping" has incompatible type "str | InvokeAgentModelInvocationOutput"; expected "str" [typeddict-item] diff --git a/tests/test_litellm/ocr/__init__.py b/tests/test_litellm/ocr/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/passthrough/__init__.py b/tests/test_litellm/passthrough/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index e42a47a1091..f014e9c26d1 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1,5 +1,6 @@ import asyncio import json +import sys import time from collections.abc import Iterator, Mapping from types import SimpleNamespace @@ -52,6 +53,7 @@ from litellm.proxy.auth.auth_checks import ( _log_budget_lookup_failure, _tag_max_budget_check, _team_max_budget_check, + _team_member_max_budget_alert_check, _virtual_key_max_budget_alert_check, _check_agent_caller_model_access, _virtual_key_max_budget_check, @@ -3774,6 +3776,141 @@ async def test_virtual_key_max_budget_alert_check_without_user_obj(): assert captured_call_info.user_email is None +@pytest.mark.parametrize( + "spend, team_metadata, expect_alert", + [ + (0.05, {"team_member_max_budget_alert_emails": {"50": [], "100": ["finance@co.com"]}}, True), + (0.10, {"team_member_max_budget_alert_emails": {"50": [], "100": ["finance@co.com"]}}, True), + (0.049, {"team_member_max_budget_alert_emails": {"50": [], "100": ["finance@co.com"]}}, False), + (0.0, {"team_member_max_budget_alert_emails": {"50": []}}, False), + (0.10, {"team_member_max_budget_alert_emails": {"abc": []}}, False), + (0.05, {"team_member_max_budget_alert_emails": {"0": ["finance@co.com"], "100": []}}, False), + (0.10, {"team_member_max_budget_alert_emails": {"101": ["finance@co.com"]}}, False), + (0.10, {"team_member_max_budget_alert_emails": "50"}, False), + (0.10, {"soft_budget_alerting_emails": ["finance@co.com"]}, False), + (0.10, None, False), + ], +) +@pytest.mark.asyncio +async def test_team_member_max_budget_alert_check_dispatches_only_at_configured_thresholds( + spend, team_metadata, expect_alert +): + captured: list[tuple[str, CallInfo]] = [] + + class RecordingProxyLogging: + async def budget_alerts(self, type, user_info): + captured.append((type, user_info)) + + _team_member_max_budget_alert_check( + team_id="team-1", + team_alias="platform", + team_metadata=team_metadata, + organization_id="org-1", + user_id="user-1", + user_email="member@co.com", + proxy_logging_obj=RecordingProxyLogging(), + spend=spend, + max_budget=0.10, + ) + await asyncio.sleep(0) + + if not expect_alert: + assert captured == [], captured + return + assert [type for type, _ in captured] == ["max_budget_alert"], captured + call_info = captured[0][1] + assert call_info.event_group == Litellm_EntityType.TEAM_MEMBER + assert (call_info.spend, call_info.max_budget) == (spend, 0.10) + assert (call_info.user_id, call_info.user_email) == ("user-1", "member@co.com") + assert (call_info.team_id, call_info.team_alias, call_info.organization_id) == ("team-1", "platform", "org-1") + assert call_info.max_budget_alert_emails == {"50": [], "100": ["finance@co.com"]} + assert call_info.token is None + + +@pytest.mark.asyncio +async def test_team_member_max_budget_alert_check_drops_thresholds_outside_1_to_100(): + captured: list[CallInfo] = [] + + class RecordingProxyLogging: + async def budget_alerts(self, type, user_info): + captured.append(user_info) + + _team_member_max_budget_alert_check( + team_id="team-1", + team_alias="platform", + team_metadata={ + "team_member_max_budget_alert_emails": { + "0": ["a@co.com"], + "50": [], + "150": ["b@co.com"], + "1" * (sys.int_info.default_max_str_digits + 1): ["c@co.com"], + } + }, + organization_id="org-1", + user_id="user-1", + user_email="member@co.com", + proxy_logging_obj=RecordingProxyLogging(), + spend=0.05, + max_budget=0.10, + ) + await asyncio.sleep(0) + + assert [call_info.max_budget_alert_emails for call_info in captured] == [{"50": []}], captured + + +@pytest.mark.asyncio +async def test_check_team_member_budget_dispatches_the_configured_alert_before_the_hard_cap(): + from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership + + captured: list[tuple[str, CallInfo]] = [] + + class RecordingProxyLogging: + async def budget_alerts(self, type, user_info): + captured.append((type, user_info)) + + team_object = LiteLLM_TeamTable( + team_id="team-1", + team_alias="platform", + metadata={"team_member_max_budget_alert_emails": {"50": [], "100": ["finance@co.com"]}}, + ) + user_object = LiteLLM_UserTable(user_id="user-1", user_email="member@co.com") + valid_token = UserAPIKeyAuth(token="tok-1", user_id="user-1", team_id="team-1") + team_membership = LiteLLM_TeamMembership( + user_id="user-1", + team_id="team-1", + spend=0.10, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10), + ) + + async def spend_from_fallback(counter_key, fallback_spend, max_budget=None, **kwargs): + return fallback_spend + + with ( + patch("litellm.proxy.proxy_server.get_current_spend", spend_from_fallback), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", new_callable=AsyncMock, return_value=team_membership + ), + ): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _check_team_member_budget( + team_object=team_object, + user_object=user_object, + valid_token=valid_token, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + proxy_logging_obj=RecordingProxyLogging(), + ) + await asyncio.sleep(0) + + assert (exc_info.value.entity_type, exc_info.value.entity_id) == ("team_member", "user-1:team-1") + assert [type for type, _ in captured] == ["max_budget_alert"], captured + call_info = captured[0][1] + assert call_info.event_group == Litellm_EntityType.TEAM_MEMBER + assert (call_info.spend, call_info.max_budget) == (0.10, 0.10) + assert (call_info.user_id, call_info.user_email, call_info.team_id) == ("user-1", "member@co.com", "team-1") + assert call_info.max_budget_alert_emails == {"50": [], "100": ["finance@co.com"]} + + @pytest.mark.parametrize( "spend, max_budget, expect_alert", [ diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index de669449f85..470db99108a 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -29,6 +29,7 @@ from litellm.proxy._types import ( LiteLLM_OrganizationTable, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, + Litellm_EntityType, LitellmUserRoles, ProxyErrorTypes, ProxyException, @@ -8055,6 +8056,152 @@ async def test_cached_key_team_member_budget_honours_temp_increase(expiry_offset assert "Max budget: 2.0" in exc_info.value.message +async def _authenticate_and_authorize(mock_request, api_key): + """Builder then the single common_checks gate, the same sequence user_api_key_auth runs.""" + from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request + + request_data = {"model": "claude-sonnet-5", "messages": [{"role": "user", "content": "hi"}]} + auth_obj = await _user_api_key_auth_builder( + request=mock_request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data=request_data, + ) + recovered = await _authorize_authenticated_request( + user_api_key_auth_obj=auth_obj, + request=mock_request, + request_data=request_data, + route="/v1/messages", + api_key=f"Bearer {api_key}", + ) + return recovered or auth_obj + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "team_member_spend, expect_blocked, expected_alerts", + [ + (1.1, False, 0), + (1.2, False, 1), + (2.4, True, 1), + ], +) +async def test_cached_key_team_member_budget_emails_configured_thresholds( + team_member_spend, expect_blocked, expected_alerts +): + """The team's team_member_max_budget_alert_emails thresholds fire from the cached-key auth path, + including on the request that trips the hard cap, and stay silent below the lowest threshold.""" + from litellm.proxy._types import LiteLLM_TeamMembership, LiteLLM_TeamTableCachedObj + from litellm.proxy.common_utils.user_api_key_cache import ( + team_membership_auth_cache_key, + team_membership_reservation_cache_key, + ) + from litellm.proxy.utils import hash_token + + api_key = "sk-team-member-alert-thresholds" + hashed_token = hash_token(api_key) + team_id = "team-alert-thresholds" + user_id = "user-alert-thresholds" + alert_emails = {"50": [], "100": ["finance@example.com"]} + + user_api_key_cache = DualCache() + await _cache_key_object( + hashed_token=hashed_token, + user_api_key_obj=UserAPIKeyAuth( + token=hashed_token, + team_id=team_id, + team_alias="platform", + team_metadata={"team_member_max_budget_alert_emails": alert_emails}, + user_id=user_id, + team_member_spend=team_member_spend, + ), + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=None, + ) + await user_api_key_cache.async_set_cache( + key=f"team_id:{team_id}", + value=LiteLLM_TeamTableCachedObj( + team_id=team_id, + team_alias="platform", + metadata={"team_member_max_budget_alert_emails": alert_emails}, + ), + ) + await user_api_key_cache.async_set_cache( + key=user_id, + value=LiteLLM_UserTable( + user_id=user_id, user_email="member@example.com", user_role=LitellmUserRoles.INTERNAL_USER + ), + ) + membership = LiteLLM_TeamMembership( + user_id=user_id, + team_id=team_id, + spend=team_member_spend, + budget_id="budget-alert-thresholds", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=2.4), + ) + # A live proxy holds the row under both keys, so any second team-member check in the + # auth flow would find it too and send a duplicate alert. + for membership_cache_key in ( + team_membership_reservation_cache_key(team_id=team_id, user_id=user_id), + team_membership_auth_cache_key(team_id=team_id, user_id=user_id), + ): + await user_api_key_cache.async_set_cache(key=membership_cache_key, value=membership) + + mock_request = MagicMock() + mock_request.url.path = "/v1/messages" + mock_request.method = "POST" + mock_request.headers = {"authorization": f"Bearer {api_key}"} + mock_request.query_params = {} + mock_request.state = SimpleNamespace() + + proxy_logging_obj = MagicMock() + proxy_logging_obj.budget_alerts = AsyncMock() + proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock(return_value=None) + + async def _auth(): + return await _authenticate_and_authorize(mock_request, api_key) + + with ( + patch( # test-quality-ok: the builder reads proxy settings from module globals, no injection seam + "litellm.proxy.proxy_server.general_settings", {"disable_budget_reservation": True} + ), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), # test-quality-ok: module-global proxy state + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: module-global proxy state + patch( # test-quality-ok: seed the cached key, team and membership without a DB + "litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache + ), + patch( # test-quality-ok: module-global proxy state + "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj + ), + patch( # test-quality-ok: the live counter needs Redis or a DB; pin the spend the check compares + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=team_member_spend), + ), + ): + if expect_blocked: + with pytest.raises(ProxyException) as exc_info: + await _auth() + assert exc_info.value.type == ProxyErrorTypes.budget_exceeded + else: + await _auth() + await asyncio.sleep(0) + + assert proxy_logging_obj.budget_alerts.await_count == expected_alerts + if expected_alerts == 0: + return + call_info = proxy_logging_obj.budget_alerts.await_args.kwargs["user_info"] + assert proxy_logging_obj.budget_alerts.await_args.kwargs["type"] == "max_budget_alert" + assert call_info.event_group == Litellm_EntityType.TEAM_MEMBER + assert (call_info.spend, call_info.max_budget) == (team_member_spend, 2.4) + assert (call_info.user_id, call_info.user_email) == (user_id, "member@example.com") + assert (call_info.team_id, call_info.team_alias) == (team_id, "platform") + assert call_info.max_budget_alert_emails == alert_emails + + async def _proxy_exception_for_key( api_key: str, general_settings: dict[str, bool], diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py deleted file mode 100644 index 78728d6fd58..00000000000 --- a/tests/test_litellm/test_main.py +++ /dev/null @@ -1,164 +0,0 @@ -import json -import os - -import pytest - - -from unittest.mock import MagicMock, patch - -import litellm - - -async def _async_fake_bedrock_image_details(image_url): - return "ZmFrZS1pbWFnZQ==", "image/png" - - -@pytest.fixture(autouse=True) -def clear_client_cache(): - """ - Clear the HTTP client cache before each test to ensure mocks are used. - This prevents cached real clients from being reused across tests. - """ - cache = getattr(litellm, "in_memory_llm_clients_cache", None) - if cache is not None: - cache.flush_cache() - yield - if cache is not None: - cache.flush_cache() - - -@pytest.fixture(autouse=True) -def add_api_keys_to_env(monkeypatch): - monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-api03-1234567890") - monkeypatch.setenv("OPENAI_API_KEY", "sk-openai-api03-1234567890") - monkeypatch.setenv("AWS_ACCESS_KEY_ID", "my-fake-aws-access-key-id") - monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "my-fake-aws-secret-access-key") - monkeypatch.setenv("AWS_REGION", "us-east-1") - # Keep these transformation tests on the simple access-key path. A leaked - # session token or role/web-identity env var pushes Bedrock auth down a - # different branch and fails before the mocked HTTP client is exercised. - monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False) - monkeypatch.delenv("AWS_ROLE_ARN", raising=False) - monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False) - - -@pytest.mark.parametrize( - "model", - [ - "gemini/gemini-1.5-flash", - "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - "bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0", - "anthropic/claude-3-5-sonnet", - ], -) -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_url_with_format_param(model, sync_mode, monkeypatch): - from litellm import acompletion, completion - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler - from litellm.litellm_core_utils.prompt_templates import factory as prompt_factory - - if sync_mode: - client = HTTPHandler() - else: - client = AsyncHTTPHandler() - - # This test is about request shaping, not live image downloads. Stub the - # URL->image conversion helpers so suite-level network/client state from - # earlier tests cannot prevent the mocked provider client from being hit. - fake_base64_image = "data:image/png;base64,ZmFrZS1pbWFnZQ==" - monkeypatch.setattr( - prompt_factory, "convert_url_to_base64", lambda url: fake_base64_image - ) - monkeypatch.setattr( - prompt_factory.BedrockImageProcessor, - "get_image_details", - staticmethod(lambda image_url: ("ZmFrZS1pbWFnZQ==", "image/png")), - ) - monkeypatch.setattr( - prompt_factory.BedrockImageProcessor, - "get_image_details_async", - staticmethod(_async_fake_bedrock_image_details), - ) - - args = { - "model": model, - "messages": [ - { - "role": "user", - "content": [ - { - "type": "image_url", - "image_url": { - "url": "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png", - "format": "image/png", - }, - }, - {"type": "text", "text": "Describe this image"}, - ], - } - ], - } - if model.startswith("gemini/"): - args["api_key"] = "test-api-key" - with patch.object(client, "post", new=MagicMock()) as mock_client: - try: - if sync_mode: - response = completion(**args, client=client) - else: - response = await acompletion(**args, client=client) - print(response) - except Exception as e: - pass - - mock_client.assert_called() - - print(mock_client.call_args.kwargs) - - if "data" in mock_client.call_args.kwargs: - json_str = mock_client.call_args.kwargs["data"] - else: - json_str = json.dumps(mock_client.call_args.kwargs["json"]) - - if isinstance(json_str, bytes): - json_str = json_str.decode("utf-8") - - print(f"type of json_str: {type(json_str)}") - - # Bedrock models convert URLs to base64, while direct Anthropic models support URLs - # bedrock/invoke models use Anthropic messages API which supports URLs - if model.startswith("bedrock/invoke/"): - # bedrock/invoke should convert URLs to base64 (doesn't support URL references) - # URL should NOT be in the JSON (it should be converted to base64) - assert "https://awsmp-logos.s3.amazonaws.com" not in json_str - # Should have base64 data in the source (type="base64", not type="url") - assert '"type":"base64"' in json_str or '"type": "base64"' in json_str - # Should have "data" field containing base64 content - assert '"data"' in json_str - elif model.startswith("bedrock/"): - # Regular Bedrock models should convert URLs to base64 (uses "bytes" field) - # URL should NOT be in the JSON (it should be converted to base64) - assert "https://awsmp-logos.s3.amazonaws.com" not in json_str - # Should have "bytes" field (Bedrock uses "bytes" not "base64" in the field name) - assert '"bytes"' in json_str or '"bytes":' in json_str - elif model.startswith("anthropic/"): - # Direct Anthropic models should pass HTTPS URLs directly (HTTP URLs are converted to base64) - # Since we're using HTTPS URL, it should be passed as-is - assert "https://awsmp-logos.s3.amazonaws.com" in json_str - # For Anthropic, URL references use "url" type, not base64 - assert '"type":"url"' in json_str or '"type": "url"' in json_str - else: - # For other models, check format parameter is respected - assert "png" in json_str - assert "jpeg" not in json_str - - -@pytest.fixture(autouse=True) -def set_openrouter_api_key(): - original_api_key = os.environ.get("OPENROUTER_API_KEY") - os.environ["OPENROUTER_API_KEY"] = "fake-key-for-testing" - yield - if original_api_key is not None: - os.environ["OPENROUTER_API_KEY"] = original_api_key - else: - del os.environ["OPENROUTER_API_KEY"] diff --git a/tests/test_rust_python_harness.py b/tests/test_rust_python_harness.py index a1bb370a074..9af38941684 100644 --- a/tests/test_rust_python_harness.py +++ b/tests/test_rust_python_harness.py @@ -36,8 +36,6 @@ def _case(module: str = "tests.example") -> HarnessCase: @pytest.mark.parametrize( "module", [ - "tests.rust-python-harness.strategies.e2e_parity.sdk.ocr.test_sdk_parity", - "tests.rust-python-harness.strategies.trace_parity.sdk.ocr.case", "tests.rust-python-harness.strategies.trace_parity.sdk.messages.case", "tests.rust-python-harness.strategies.trace_parity.sdk.chat_completions.case", "tests.rust-python-harness.strategies.trace_parity.sdk.transcription.case", diff --git a/tests/unit/AGENTS.md b/tests/unit/AGENTS.md index 191777f3e83..fc9798ad30e 100644 --- a/tests/unit/AGENTS.md +++ b/tests/unit/AGENTS.md @@ -30,7 +30,7 @@ Green if `send_batched` drops every row. pydantic doubles in 12 of 203 files, fa ## Where it goes `tests/unit/` mirrors `litellm/`, so a changed file selects its tests by path, not a mapping -file. Empty today; new unit tests go here. The examples above live in `tests/test_litellm` +file. New unit tests go here ## Writing it so a human can read it diff --git a/tests/unit/a2a_protocol/test_cost_calculator.py b/tests/unit/a2a_protocol/test_cost_calculator.py index 56d3d57c89e..8d8ec815f3a 100644 --- a/tests/unit/a2a_protocol/test_cost_calculator.py +++ b/tests/unit/a2a_protocol/test_cost_calculator.py @@ -10,6 +10,12 @@ import pytest import litellm from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + + +async def _reset_callbacks_and_settle_pending_logs() -> None: + litellm.logging_callback_manager._reset_all_callbacks() + await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0) def _make_send_message_request(request_id: str, user_text: str = "Hello"): @@ -129,7 +135,7 @@ async def test_asend_message_uses_cost_per_query(monkeypatch): from litellm.a2a_protocol import asend_message # Setup logger - litellm.logging_callback_manager._reset_all_callbacks() + await _reset_callbacks_and_settle_pending_logs() cost_logger = CostLogger() monkeypatch.setattr(litellm, "callbacks", [cost_logger]) @@ -164,7 +170,7 @@ async def test_asend_message_uses_cost_per_query_from_litellm_params_dict(monkey """ from litellm.a2a_protocol import asend_message - litellm.logging_callback_manager._reset_all_callbacks() + await _reset_callbacks_and_settle_pending_logs() cost_logger = CostLogger() monkeypatch.setattr(litellm, "callbacks", [cost_logger]) @@ -225,7 +231,7 @@ async def test_asend_message_uses_input_output_cost_per_token(monkeypatch): from litellm.a2a_protocol import asend_message # Setup logger - litellm.logging_callback_manager._reset_all_callbacks() + await _reset_callbacks_and_settle_pending_logs() token_cost_logger = TokenAndCostLogger() monkeypatch.setattr(litellm, "callbacks", [token_cost_logger]) @@ -299,7 +305,7 @@ async def test_asend_message_passes_agent_id_to_callback(monkeypatch): from litellm.a2a_protocol import asend_message # Setup logger - litellm.logging_callback_manager._reset_all_callbacks() + await _reset_callbacks_and_settle_pending_logs() agent_id_logger = AgentIdLogger() monkeypatch.setattr(litellm, "callbacks", [agent_id_logger]) @@ -359,7 +365,7 @@ async def test_asend_message_streaming_propagates_metadata(): from litellm.a2a_protocol import asend_message_streaming # Setup logger - litellm.logging_callback_manager._reset_all_callbacks() + await _reset_callbacks_and_settle_pending_logs() metadata_logger = MetadataLogger() litellm.logging_callback_manager.add_litellm_async_success_callback(metadata_logger) @@ -406,7 +412,7 @@ async def test_asend_message_streaming_triggers_callbacks(): from litellm.a2a_protocol import asend_message_streaming # Setup logger - must use logging_callback_manager to properly register - litellm.logging_callback_manager._reset_all_callbacks() + await _reset_callbacks_and_settle_pending_logs() callback_logger = AgentIdLogger() litellm.logging_callback_manager.add_litellm_async_success_callback(callback_logger) litellm.logging_callback_manager.add_litellm_success_callback(callback_logger) diff --git a/tests/unit/caching/test_redis_cache.py b/tests/unit/caching/test_redis_cache.py index 5d72fe7213d..5f83be7c7bc 100644 --- a/tests/unit/caching/test_redis_cache.py +++ b/tests/unit/caching/test_redis_cache.py @@ -2,6 +2,7 @@ import asyncio import time from collections.abc import Iterator from datetime import timedelta +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -1023,7 +1024,12 @@ async def test_breaker_metrics_track_state_and_failure_class(): from redis.exceptions import ConnectionError as RedisConnectionError from redis.exceptions import TimeoutError as RedisTimeoutError - from litellm.caching.redis_cache import RedisCircuitBreaker, is_redis_timeout_failure + from litellm.caching.redis_cache import RedisCircuitBreaker, _breaker_metrics, is_redis_timeout_failure + + metrics: Final = _breaker_metrics() + for collector in (metrics._state_gauge, metrics._transitions, metrics._failures): + if collector is not None and collector not in REGISTRY._collector_to_names: + REGISTRY.register(collector) def sample(name, labels=None): return REGISTRY.get_sample_value(name, labels) or 0.0 diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py index ecea4723bf4..ec957d80904 100644 --- a/tests/unit/conftest.py +++ b/tests/unit/conftest.py @@ -12,6 +12,32 @@ import httpx import pytest from pytest_socket import enable_socket, socket_allow_hosts +HOST_ENVIRONMENT_ALLOWLIST: Final = frozenset( + ( + "PATH", + "HOME", + "USER", + "LOGNAME", + "TMPDIR", + "TEMP", + "TMP", + "LANG", + "LC_ALL", + "LC_CTYPE", + "TZ", + "VIRTUAL_ENV", + "LITELLM_LOCAL_MODEL_COST_MAP", + "TIKTOKEN_CACHE_DIR", + ) +) +HOST_ENVIRONMENT_ALLOWED_PREFIXES: Final = ("PYTEST_", "PYTHON", "COV_CORE_", "COVERAGE_") +HOST_ONLY_ENVIRONMENT: Final = frozenset( + name + for name in os.environ + if name not in HOST_ENVIRONMENT_ALLOWLIST and not name.startswith(HOST_ENVIRONMENT_ALLOWED_PREFIXES) +) + +os.environ["PYTHON_DOTENV_DISABLED"] = "1" os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" import litellm # noqa: E402 # litellm reads LITELLM_LOCAL_MODEL_COST_MAP at import @@ -170,6 +196,8 @@ def isolated_aws_config_files(tmp_path_factory: pytest.TempPathFactory) -> tuple def isolate_host_environment(isolated_aws_config_files: tuple[Path, Path]) -> Iterator[None]: credentials, config = isolated_aws_config_files with pytest.MonkeyPatch.context() as environment: + for name in HOST_ONLY_ENVIRONMENT: + environment.delenv(name, raising=False) environment.setenv("AWS_SHARED_CREDENTIALS_FILE", str(credentials)) environment.setenv("AWS_CONFIG_FILE", str(config)) environment.setenv("AWS_EC2_METADATA_DISABLED", "true") diff --git a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py index 8b89c592f02..52e44ca5448 100644 --- a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py +++ b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py @@ -1090,6 +1090,47 @@ async def test_multi_threshold_empty_emails_only_owner( assert to_emails == ["owner@co.com"] +@pytest.mark.asyncio +async def test_multi_threshold_team_member_alert_renders_member_template_per_team( + base_email_logger, mock_send_email +): + """A team member budget alert is keyed per member and team, names the member and team, + and goes to the member plus the threshold's configured recipients""" + user_info = CallInfo( + user_id="member_1", + user_email="member@co.com", + team_id="team_a", + team_alias="Platform", + spend=0.10, + max_budget=0.10, + event_group=Litellm_EntityType.TEAM_MEMBER, + max_budget_alert_emails={"50": [], "100": ["finance@co.com"]}, + ) + + mock_cache = mock.AsyncMock() + mock_cache.async_increment_cache = mock.AsyncMock(return_value=1) + base_email_logger.internal_usage_cache = mock_cache + + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info) + + cache_keys = sorted(c[1]["key"] for c in mock_cache.async_increment_cache.call_args_list) + assert cache_keys == [ + "email_budget_alerts:max_budget_alert:100:team_member:member_1:team_a", + "email_budget_alerts:max_budget_alert:50:team_member:member_1:team_a", + ] + assert mock_send_email.call_count == 2 + hundred = next( + c.kwargs for c in mock_send_email.call_args_list if "100%" in c.kwargs["subject"] + ) + assert hundred["subject"] == "LiteLLM: Team Member Budget Alert - 100% of Team Member Budget Reached" + assert sorted(hundred["to_email"]) == ["finance@co.com", "member@co.com"] + assert "member@co.com" in hundred["html_body"] and "Platform" in hundred["html_body"] + assert "team member budget" in hundred["html_body"] and "$0.1" in hundred["html_body"] + fifty = next(c.kwargs for c in mock_send_email.call_args_list if "50%" in c.kwargs["subject"]) + assert fifty["to_email"] == ["member@co.com"] + + @pytest.mark.asyncio async def test_no_map_preserves_old_single_threshold( base_email_logger, mock_send_email diff --git a/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index 92ff3d5813c..f1c80bb11ea 100644 --- a/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -1,7 +1,6 @@ import asyncio -import logging from datetime import datetime, timedelta, timezone from unittest.mock import MagicMock, call, patch @@ -9,7 +8,6 @@ import pytest from prometheus_client import REGISTRY import litellm -from litellm._logging import verbose_logger from litellm.types.utils import ( StandardLoggingHiddenParams, StandardLoggingMetadata, @@ -27,10 +25,6 @@ except Exception: PrometheusLogger = None from litellm.proxy._types import UserAPIKeyAuth -verbose_logger.setLevel(logging.DEBUG) - -litellm.set_verbose = True - @pytest.fixture def prometheus_logger() -> PrometheusLogger: diff --git a/tests/unit/enterprise/integrations/test_prometheus.py b/tests/unit/enterprise/integrations/test_prometheus.py index 7315f2b9881..16d6ff9d9a0 100644 --- a/tests/unit/enterprise/integrations/test_prometheus.py +++ b/tests/unit/enterprise/integrations/test_prometheus.py @@ -477,7 +477,7 @@ def test_valid_configuration_passes_validation(): # ============================================================================== -@pytest.fixture +@pytest.fixture(autouse=True) def reset_prometheus_exclude_settings(): """Restore the global exclude settings after each test so they don't leak.""" prev_metrics = litellm.prometheus_exclude_metrics diff --git a/tests/unit/integrations/SlackAlerting/test_budget_alert_types.py b/tests/unit/integrations/SlackAlerting/test_budget_alert_types.py index 52b7cc983a7..f3199d9ebf9 100644 --- a/tests/unit/integrations/SlackAlerting/test_budget_alert_types.py +++ b/tests/unit/integrations/SlackAlerting/test_budget_alert_types.py @@ -1,4 +1,7 @@ -from litellm.integrations.SlackAlerting.budget_alert_types import SoftBudgetAlert +from litellm.integrations.SlackAlerting.budget_alert_types import ( + SoftBudgetAlert, + TokenBudgetAlert, +) from litellm.proxy._types import CallInfo, Litellm_EntityType @@ -64,3 +67,31 @@ class TestSoftBudgetAlert: result = alert.get_id(user_info) assert result == "default_id" + + +class TestTokenBudgetAlert: + def test_get_id_dedupes_team_member_alerts_per_member_and_team(self): + alert = TokenBudgetAlert() + team_a = CallInfo( + spend=8.0, max_budget=10.0, user_id="member_1", team_id="team_a", event_group=Litellm_EntityType.TEAM_MEMBER + ) + team_b = CallInfo( + spend=8.0, max_budget=10.0, user_id="member_1", team_id="team_b", event_group=Litellm_EntityType.TEAM_MEMBER + ) + + assert alert.get_id(team_a) == "team_member:member_1:team_a" + assert alert.get_id(team_b) == "team_member:member_1:team_b" + + def test_get_id_uses_token_for_key_alerts(self): + alert = TokenBudgetAlert() + user_info = CallInfo( + spend=8.0, + max_budget=10.0, + token="hashed_key", + user_id="member_1", + team_id="team_a", + event_group=Litellm_EntityType.KEY, + ) + + assert alert.get_id(user_info) == "hashed_key" + assert alert.get_event_message() == "Key Budget: " diff --git a/tests/unit/integrations/SlackAlerting/test_slack_alerting.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting.py index b9e5ff2eeb7..0c2b95fd448 100644 --- a/tests/unit/integrations/SlackAlerting/test_slack_alerting.py +++ b/tests/unit/integrations/SlackAlerting/test_slack_alerting.py @@ -393,6 +393,33 @@ def _slack_alerting_with_env_resolution() -> SlackAlerting: return slack_alerting +@pytest.mark.asyncio +@pytest.mark.parametrize( + "event_group, expected_prefix", + [ + (Litellm_EntityType.TEAM_MEMBER, "Team Member Budget: Budget Crossed"), + (Litellm_EntityType.KEY, "Key Budget: Budget Crossed"), + ], +) +async def test_max_budget_alert_labels_team_member_budget(event_group, expected_prefix): + slack_alerting: Final = _slack_alerting_with_env_resolution() + slack_alerting.send_alert = AsyncMock() + + await slack_alerting.budget_alerts( + type="max_budget_alert", + user_info=CallInfo( + spend=10.5, + max_budget=10.0, + token="hashed_key", + user_id="member_1", + team_id="team_a", + event_group=event_group, + ), + ) + + assert slack_alerting.send_alert.await_args.kwargs["message"].startswith(expected_prefix) + + @pytest.mark.asyncio async def test_send_alert_falls_back_to_alerting_webhook_url_env(monkeypatch): monkeypatch.delenv("SLACK_WEBHOOK_URL", raising=False) diff --git a/tests/unit/integrations/dotprompt/test_prompt_manager.py b/tests/unit/integrations/dotprompt/test_prompt_manager.py index 51e14b61929..dbd4e4a4c4c 100644 --- a/tests/unit/integrations/dotprompt/test_prompt_manager.py +++ b/tests/unit/integrations/dotprompt/test_prompt_manager.py @@ -22,7 +22,7 @@ def test_prompt_manager_initialization(): # Test with the existing prompts directory prompt_dir = Path( __file__ - ).parent # Current directory when running from tests/test_litellm/prompts + ).parent manager = PromptManager(prompt_directory=str(prompt_dir)) # Should have loaded at least the sample prompts @@ -56,7 +56,7 @@ def test_render_simple_template(): """Test rendering a simple template with variables.""" prompt_dir = Path( __file__ - ).parent # Current directory when running from tests/test_litellm/prompts + ).parent manager = PromptManager(prompt_directory=str(prompt_dir)) # Test sample_prompt rendering @@ -72,7 +72,7 @@ def test_render_chat_prompt(): """Test rendering the chat prompt with conditional content.""" prompt_dir = Path( __file__ - ).parent # Current directory when running from tests/test_litellm/prompts + ).parent manager = PromptManager(prompt_directory=str(prompt_dir)) # Test with system context @@ -98,7 +98,7 @@ def test_render_coding_assistant(): """Test rendering the coding assistant prompt with complex logic.""" prompt_dir = Path( __file__ - ).parent # Current directory when running from tests/test_litellm/prompts + ).parent manager = PromptManager(prompt_directory=str(prompt_dir)) rendered = manager.render( @@ -159,7 +159,7 @@ def test_prompt_not_found(): """Test error handling for non-existent prompts.""" prompt_dir = Path( __file__ - ).parent # Current directory when running from tests/test_litellm/prompts + ).parent manager = PromptManager(prompt_directory=str(prompt_dir)) with pytest.raises(KeyError, match="Prompt 'nonexistent' not found"): @@ -170,7 +170,7 @@ def test_list_prompts(): """Test listing available prompts.""" prompt_dir = Path( __file__ - ).parent # Current directory when running from tests/test_litellm/prompts + ).parent manager = PromptManager(prompt_directory=str(prompt_dir)) prompts = manager.list_prompts() @@ -184,7 +184,7 @@ def test_get_prompt_metadata(): """Test retrieving prompt metadata.""" prompt_dir = Path( __file__ - ).parent # Current directory when running from tests/test_litellm/prompts + ).parent manager = PromptManager(prompt_directory=str(prompt_dir)) metadata = manager.get_prompt_metadata("sample_prompt") @@ -221,7 +221,7 @@ def test_add_prompt_programmatically(): """Test adding prompts programmatically.""" prompt_dir = Path( __file__ - ).parent # Current directory when running from tests/test_litellm/prompts + ).parent manager = PromptManager(prompt_directory=str(prompt_dir)) initial_count = len(manager.prompts) diff --git a/tests/unit/integrations/websearch_interception/test_websearch_chat_completion.py b/tests/unit/integrations/websearch_interception/test_websearch_chat_completion.py index 7ef43e2eadf..21e50561f57 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_chat_completion.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_chat_completion.py @@ -5,7 +5,6 @@ Tests the end-to-end flow of websearch_interception callback with litellm.acompletion() for transparent server-side web search execution. """ -import os from unittest.mock import MagicMock import pytest @@ -37,75 +36,6 @@ def websearch_logger(): return WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI, LlmProviders.MINIMAX]) -@pytest.mark.asyncio -@pytest.mark.skipif( - os.environ.get("OPENAI_API_KEY") is None, - reason="OPENAI_API_KEY not set", -) -async def test_websearch_chat_completion_with_openai(): - """Test websearch interception with OpenAI chat completions API. - - This test verifies that: - 1. Model calls litellm_web_search tool - 2. Server executes web search automatically - 3. Server makes follow-up request with search results - 4. User gets final answer without tool_calls - """ - # Configure WebSearch interception - original_callbacks = litellm.callbacks.copy() if litellm.callbacks else [] - websearch_logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI]) - litellm.callbacks = [websearch_logger] - - try: - response = await litellm.acompletion( - model="gpt-4o-mini", # Use cheaper model for testing - messages=[ - { - "role": "user", - "content": "What's the weather in San Francisco today?", - } - ], - tools=[ - { - "type": "function", - "function": { - "name": "litellm_web_search", - "description": "Search the web for information", - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "Search query", - } - }, - "required": ["query"], - }, - }, - } - ], - ) - - # Verify response structure - assert isinstance(response, ModelResponse) - assert response.choices[0].message.content is not None - assert len(response.choices[0].message.content) > 0 - - # If agentic loop worked, we should NOT have tool_calls in final response - # (they should have been executed and replaced with final answer) - if hasattr(response.choices[0].message, "tool_calls"): - # If tool_calls exist, it means agentic loop didn't run - # This could happen if search tool is not configured - pytest.skip("Agentic loop did not execute - search tool may not be configured") - - # Verify we got a meaningful response - assert response.choices[0].finish_reason in ["stop", "end_turn"] - - finally: - # Restore original callbacks - litellm.callbacks = original_callbacks - - @pytest.mark.asyncio async def test_websearch_chat_completion_hook_detection(): """Test that websearch hook correctly detects tool calls in response.""" @@ -321,61 +251,6 @@ async def test_websearch_json_serialization_fix(): assert arguments_str != "{'query': 'weather in SF'}" -@pytest.mark.asyncio -@pytest.mark.skipif( - os.environ.get("OPENAI_API_KEY") is None or os.environ.get("PERPLEXITY_API_KEY") is None, - reason="OPENAI_API_KEY or PERPLEXITY_API_KEY not set", -) -async def test_websearch_streaming_conversion(): - """Test that streaming requests are converted to non-streaming for web search. - - When stream=True is passed with web search tools, the handler should: - 1. Convert stream=True to stream=False for initial request - 2. Execute web search - 3. Convert final response back to streaming - """ - websearch_logger = WebSearchInterceptionLogger( - enabled_providers=[LlmProviders.OPENAI], search_tool_name="perplexity-search" - ) - litellm.callbacks = [websearch_logger] - - try: - response = await litellm.acompletion( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "What's the latest AI news?"}], - tools=[ - { - "type": "function", - "function": { - "name": "litellm_web_search", - "description": "Search the web", - "parameters": { - "type": "object", - "properties": {"query": {"type": "string"}}, - }, - }, - } - ], - stream=True, - ) - - # Response should be a streaming iterator - chunks = [] - async for chunk in response: - chunks.append(chunk) - - # Verify we got streaming chunks - assert len(chunks) > 0 - - # Verify chunks have expected structure - for chunk in chunks: - assert hasattr(chunk, "choices") - assert len(chunk.choices) > 0 - - finally: - litellm.callbacks = [] - - @pytest.mark.asyncio async def test_maybe_run_chat_completion_agentic_loop_calls_chat_completion_hook(): """Regression test: maybe_run_chat_completion_agentic_loop must call diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 781a3a7c4ed..0afd989272e 100644 --- a/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -3789,3 +3789,15 @@ def test_azure_gpt_6_foundry_price_sheet(_local_model_cost_map, model_base): assert azure_ai_info[field] == base assert azure_us_info[field] == pytest.approx(1.1 * base) assert azure_eu_info[field] == pytest.approx(1.2 * base) + + +@pytest.mark.parametrize("region_prefix", ["azure/", "azure/us/", "azure/eu/"]) +def test_azure_gpt_5_6_alias_matches_sol_pricing(_local_model_cost_map, region_prefix): + """The bare gpt-5.6 alias routes to GPT-5.6 Sol, so every Azure region must bill the + alias exactly like the Sol entry (including the Sept 2026 $4/$20 promo).""" + alias = litellm.model_cost[f"{region_prefix}gpt-5.6"] + sol = litellm.model_cost[f"{region_prefix}gpt-5.6-sol"] + shared_cost_fields = [f for f in alias if "cost" in f and f in sol and not isinstance(alias[f], dict)] + assert shared_cost_fields + for field in shared_cost_fields: + assert alias[field] == sol[field], field diff --git a/tests/unit/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py index c3478c0d5eb..47c4576f91f 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -1,5 +1,6 @@ """Test health check helper functions""" +import socket import struct import zlib from types import MappingProxyType @@ -214,24 +215,26 @@ async def test_ahealth_check_failure_masks_raw_request_headers(): This tests the fix for the security vulnerability where Authorization headers were being exposed in health check error responses. """ - # Use a model configuration that will fail (invalid endpoint) test_api_key = "dapi-test-key-1234567890abcdef" test_headers = { "Authorization": f"Bearer {test_api_key}", "Content-Type": "application/json", } - response = await ahealth_check( - model_params={ - "model": "databricks/dbrx-instruct", - "api_base": "https://invalid-endpoint-that-will-fail.com/", - "api_key": test_api_key, - "headers": test_headers, - }, - mode="chat", - ) + with socket.socket() as reserved: + reserved.bind(("127.0.0.1", 0)) + api_base = f"http://127.0.0.1:{reserved.getsockname()[1]}/" + + response = await ahealth_check( + model_params={ + "model": "databricks/dbrx-instruct", + "api_base": api_base, + "api_key": test_api_key, + "headers": test_headers, + }, + mode="chat", + ) - # Should have error and raw_request_typed_dict assert "error" in response assert "raw_request_typed_dict" in response @@ -243,22 +246,15 @@ async def test_ahealth_check_failure_masks_raw_request_headers(): headers = raw_request_dict["raw_request_headers"] assert headers is not None - # Security check: Authorization header should be masked, not show full key - if "Authorization" in headers: - auth_header = headers["Authorization"] - # Should be masked (e.g., "Be****90" or similar) - assert auth_header != f"Bearer {test_api_key}", "Authorization header must be masked" - assert auth_header != test_api_key, "API key must not appear in Authorization header" - # Masked headers typically have asterisks or are truncated - assert "*" in auth_header or len(auth_header) < len(f"Bearer {test_api_key}"), ( - f"Authorization header should be masked but got: {auth_header}" - ) + assert "Authorization" in headers + auth_header = headers["Authorization"] + assert auth_header != f"Bearer {test_api_key}", "Authorization header must be masked" + assert auth_header != test_api_key, "API key must not appear in Authorization header" + assert "*" in auth_header or len(auth_header) < len(f"Bearer {test_api_key}"), ( + f"Authorization header should be masked but got: {auth_header}" + ) - # Content-Type should remain unmasked (not sensitive) - if "Content-Type" in headers: - assert headers["Content-Type"] == "application/json" - - print(f"Masked Authorization header: {headers.get('Authorization', 'NOT FOUND')}") + assert headers["Content-Type"] == "application/json" @pytest.mark.asyncio diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index d717718cba2..cb1e281e356 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -7969,6 +7969,29 @@ def test_deployment_pricing_model_info_honors_a_tier_only_batch_override_over_th assert {key: info[key] for key in carried_keys} == {key: _PUBLISHED_BATCH_RATES[key] for key in carried_keys} +@pytest.mark.parametrize( + "override_key", + ( + "output_cost_per_token_above_200k_tokens_batches", + "cache_read_input_token_cost_above_200k_tokens_batches", + "cache_creation_input_token_cost_above_200k_tokens_batches", + ), +) +def test_deployment_pricing_model_info_honors_a_200k_tier_batch_override( + _published_batch_model: None, override_key: str +) -> None: + from litellm.litellm_core_utils.litellm_logging import deployment_pricing_model_info + + info: Final = deployment_pricing_model_info(_batch_deployment_id({override_key: 1e-3}), _PUBLISHED_BATCH_DEPLOYMENT) + carried_keys: Final = tuple( + key for key in (*_PUBLISHED_INPUT_BATCH_KEYS, *_PUBLISHED_OUTPUT_BATCH_KEYS) if key != override_key + ) + + assert info is not None + assert info[override_key] == 1e-3 + assert {key: info[key] for key in carried_keys} == {key: _PUBLISHED_BATCH_RATES[key] for key in carried_keys} + + def test_get_status_fields_ranks_guardrail_flagged_between_success_and_intervened(): """LIT-6894: a non-blocking flagged verdict must outrank success in the request-level guardrail_status but never mask an intervention.""" diff --git a/tests/unit/litellm_core_utils/test_logging_worker.py b/tests/unit/litellm_core_utils/test_logging_worker.py index 2bb93a58531..5d4c9e65d9b 100644 --- a/tests/unit/litellm_core_utils/test_logging_worker.py +++ b/tests/unit/litellm_core_utils/test_logging_worker.py @@ -180,6 +180,39 @@ class TestLoggingWorker: assert sorted(fired) == ["first", "second"] + def test_callback_finishing_after_loop_change_settles_only_its_own_queue(self): + worker = LoggingWorker(timeout=1.0, max_queue_size=10, concurrency=1) + fired = [] + + async def marker(name, delay=0.0): + await asyncio.sleep(delay) + fired.append(name) + + async def start_slow_callback(): + worker.ensure_initialized_and_enqueue(marker("slow", delay=0.05)) + await asyncio.sleep(0.01) + + async def log_on_second_loop(): + for name in ("b1", "b2", "b3"): + worker.ensure_initialized_and_enqueue(marker(name)) + for _ in range(2): + await asyncio.sleep(0) + + first_loop = asyncio.new_event_loop() + try: + first_loop.run_until_complete(start_slow_callback()) + first_loop_tasks = tuple(asyncio.all_tasks(first_loop)) + asyncio.run(log_on_second_loop()) + first_loop.run_until_complete(asyncio.sleep(0.1)) + failures = [ + task.exception() for task in first_loop_tasks if task.done() and not task.cancelled() and task.exception() + ] + finally: + first_loop.close() + + assert failures == [] + assert "slow" in fired + @pytest.mark.parametrize("stranded", ["still_queued", "dequeued_never_started"]) def test_flush_on_new_loop_drains_tasks_stranded_on_previous_loop(self, stranded): """ diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index b1a14e61b96..f7ded4f3fa8 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -3,16 +3,22 @@ import asyncio import base64 import importlib +import json +import os +import subprocess +import sys import threading import time -import traceback +from collections.abc import Mapping from concurrent.futures import Future, wait +from pathlib import Path from typing import Final from unittest.mock import MagicMock import anyio.to_thread import pytest import tiktoken +from tokenizers import Regex, Tokenizer, models, pre_tokenizers from unittest.mock import AsyncMock, patch @@ -23,6 +29,7 @@ import litellm.constants from litellm.constants import TOKEN_COUNTER_MAX_CONCURRENT_COUNTS from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.token_counter import ( + _encoding_count, _get_exact_count_function, _get_extrapolating_count_function, _get_tiktoken_count_function, @@ -73,15 +80,17 @@ def test_token_counter_basic(): ) -def test_token_counter_large_repeated_text_is_fast(): - messages = [{"role": "user", "content": [{"type": "text", "text": "A" * 1024 * 1024}]}] +def test_token_counter_large_repeated_text_is_encoded_in_bounded_chunks(): + text_length: Final = 1024 * 1024 + messages: Final = [{"role": "user", "content": [{"type": "text", "text": "A" * text_length}]}] - start_time = time.perf_counter() - tokens = token_counter_new(model="us.anthropic.claude-sonnet-4-6", messages=messages) - elapsed = time.perf_counter() - start_time + with patch("litellm.litellm_core_utils.token_counter._encoding_count", wraps=_encoding_count) as encoding_count: + tokens: Final = token_counter_new(model="us.anthropic.claude-sonnet-4-6", messages=messages) - assert elapsed < 2, f"Token counting took too long: {elapsed:.2f}s" + encoded_lengths: Final = tuple(len(call.args[1]) for call in encoding_count.call_args_list) assert tokens > 0 + assert sum(encoded_lengths) >= text_length + assert max(encoded_lengths) <= litellm.constants.TIKTOKEN_ENCODE_MAX_CHUNK_SIZE_CHARS @pytest.mark.parametrize( @@ -442,35 +451,46 @@ class NeedsToleranceUpdateError(Exception): # test_tokenizers() -def test_encoding_and_decoding(): - try: - sample_text = "Hellö World, this is my input string!" - # openai encoding + decoding - openai_tokens = encode(model="gpt-3.5-turbo", text=sample_text) - openai_text = decode(model="gpt-3.5-turbo", tokens=openai_tokens) +def test_encoding_and_decoding(tmp_path: Path): + sample_text = "Hellö World, this is my input string!" - assert openai_text == sample_text + # openai encoding + decoding + openai_tokens = encode(model="gpt-3.5-turbo", text=sample_text) + openai_text = decode(model="gpt-3.5-turbo", tokens=openai_tokens) - # claude encoding + decoding - claude_tokens = encode(model="claude-3-5-haiku-20241022", text=sample_text) + assert openai_text == sample_text - claude_text = decode(model="claude-3-5-haiku-20241022", tokens=claude_tokens) + # claude encoding + decoding + claude_tokens = encode(model="claude-3-5-haiku-20241022", text=sample_text) - assert claude_text == sample_text + claude_text = decode(model="claude-3-5-haiku-20241022", tokens=claude_tokens) - # cohere encoding + decoding - cohere_tokens = encode(model="command-nightly", text=sample_text) - cohere_text = decode(model="command-nightly", tokens=cohere_tokens) + assert claude_text == sample_text - assert cohere_text == sample_text + # cohere encoding + decoding + cohere_tokens = encode(model="command-nightly", text=sample_text) + cohere_text = decode(model="command-nightly", tokens=cohere_tokens) - # llama2 encoding + decoding - llama2_tokens = encode(model="meta-llama/Llama-2-7b-chat", text=sample_text) - llama2_text = decode(model="meta-llama/Llama-2-7b-chat", tokens=llama2_tokens) + assert cohere_text == sample_text - assert llama2_text == sample_text - except Exception as e: - pytest.fail(f"An exception occured: {e}\n{traceback.format_exc()}") + # llama2 encoding + decoding + words = sample_text.split() + result = _run_in_memory_hub( + HUB_ROUND_TRIP_SCRIPT, + { + "hf-internal-testing/llama-tokenizer": _word_level_tokenizer_json( + pre_tokenizers.WhitespaceSplit(), + vocab={"[UNK]": 0, **{word: i + 1 for i, word in enumerate(words)}}, + ) + }, + sample_text, + tmp_path, + ) + + assert result["decoded"] == sample_text + assert result["requested"] == ["hf-internal-testing/llama-tokenizer"] + assert len(result["tokens"]) == len(words) + assert len(result["tokens"]) != len(encode(model="gpt-3.5-turbo", text=sample_text)) # test_encoding_and_decoding() @@ -1439,3 +1459,105 @@ def test_high_detail_image_token_upper_bound_covers_every_image_size(width: int, def test_high_detail_image_token_upper_bound_is_reached_by_the_largest_high_res_image() -> None: assert calculate_img_tokens(_png_data_url(2000, 768), mode="high") == high_detail_image_token_upper_bound() assert calculate_img_tokens(_png_data_url(1, 1), mode="high") < high_detail_image_token_upper_bound() + + +HUB_SETUP_SCRIPT: Final = """ +import json +import sys +sys.path.insert(0, sys.argv[1]) +import httpx +import huggingface_hub +import litellm +served = json.loads(sys.argv[2]) +text = sys.argv[3] +requested = [] +def handle(request): + repo = request.url.path.lstrip("/").split("/resolve/")[0] + if repo not in served or not request.url.path.endswith("/tokenizer.json"): + return httpx.Response(404) + requested.append(repo) + payload = served[repo].encode() + headers = {"content-length": str(len(payload)), "etag": '"fixture"', "x-repo-commit": "a" * 40} + return httpx.Response(200, headers=headers, content=payload if request.method == "GET" else b"") +huggingface_hub.set_client_factory(lambda: httpx.Client(transport=httpx.MockTransport(handle))) +""" + +HUB_TOKENIZER_SCRIPT: Final = HUB_SETUP_SCRIPT + """ +litellm.cohere_models = {"command-r-v1"} +litellm.anthropic_models = {"claude-2"} +custom = litellm.create_pretrained_tokenizer("Xenova/llama-3-tokenizer") +print(json.dumps({ + "llama2": litellm.token_counter(model="meta-llama/Llama-2-7b-chat", text=text), + "llama3": litellm.token_counter(model="meta-llama/llama-3-70b-instruct", text=text), + "cohere": litellm.token_counter(model="command-r-v1", text=text), + "anthropic": litellm.token_counter(model="claude-2", text=text), + "custom": litellm.token_counter(custom_tokenizer=custom, text=text), + "requested": sorted(set(requested)), +})) +""" + +HUB_ROUND_TRIP_SCRIPT: Final = HUB_SETUP_SCRIPT + """ +tokens = litellm.encode(model="meta-llama/Llama-2-7b-chat", text=text) +print(json.dumps({"tokens": tokens, "decoded": litellm.decode(model="meta-llama/Llama-2-7b-chat", tokens=tokens), "requested": sorted(set(requested))})) +""" + + +def _word_level_tokenizer_json( + pre_tokenizer: pre_tokenizers.PreTokenizer, vocab: Mapping[str, int] | None = None +) -> str: + tokenizer: Final = Tokenizer( + models.WordLevel(vocab=dict(vocab) if vocab is not None else {"[UNK]": 0}, unk_token="[UNK]") + ) + tokenizer.pre_tokenizer = pre_tokenizer + return tokenizer.to_str() + + +def _run_in_memory_hub(script: str, served: dict[str, str], text: str, tmp_path: Path) -> dict: + result: Final = subprocess.run( + [ + sys.executable, + "-I", + "-c", + script, + str(Path(litellm.__file__).parent.parent), + json.dumps(served), + text, + ], + capture_output=True, + text=True, + timeout=60, + env={ + **os.environ, + "HF_HOME": str(tmp_path / "home"), + "HF_HUB_CACHE": str(tmp_path / "cache"), + "HF_ENDPOINT": "http://127.0.0.1:9", + "HF_HUB_OFFLINE": "0", + "LITELLM_LOCAL_MODEL_COST_MAP": "True", + }, + ) + + assert result.returncode == 0, result.stdout + result.stderr + return json.loads(result.stdout.strip().splitlines()[-1]) + + +def test_token_counter_uses_the_tokenizer_of_each_model_family_and_of_a_custom_tokenizer(tmp_path: Path) -> None: + sample: Final = "Tokenizers disagree: anthropic, tiktoken; llama-2 & llama-3!" + served: Final = { + "hf-internal-testing/llama-tokenizer": _word_level_tokenizer_json(pre_tokenizers.WhitespaceSplit()), + "Xenova/llama-3-tokenizer": _word_level_tokenizer_json(pre_tokenizers.Split(Regex("."), "isolated")), + "Xenova/c4ai-command-r-v01-tokenizer": _word_level_tokenizer_json(pre_tokenizers.Whitespace()), + } + expected: Final = {repo: len(Tokenizer.from_str(payload).encode(sample).ids) for repo, payload in served.items()} + anthropic_count: Final = len(Tokenizer.from_str(claude_json_str).encode(sample).ids) + tiktoken_count: Final = litellm.token_counter(model="gpt-3.5-turbo", text=sample) + assert len({*expected.values(), anthropic_count, tiktoken_count}) == len(expected) + 2 + + counts: Final = _run_in_memory_hub(HUB_TOKENIZER_SCRIPT, served, sample, tmp_path) + assert counts == { + "llama2": expected["hf-internal-testing/llama-tokenizer"], + "llama3": expected["Xenova/llama-3-tokenizer"], + "cohere": expected["Xenova/c4ai-command-r-v01-tokenizer"], + "anthropic": anthropic_count, + "custom": expected["Xenova/llama-3-tokenizer"], + "requested": sorted(served), + } diff --git a/tests/unit/litellm_core_utils/test_tokenizer.py b/tests/unit/litellm_core_utils/test_tokenizer.py index a9005ff6a86..9d08442b164 100644 --- a/tests/unit/litellm_core_utils/test_tokenizer.py +++ b/tests/unit/litellm_core_utils/test_tokenizer.py @@ -17,17 +17,13 @@ from litellm.utils import claude_json_str from tests.unit.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON -OFFLINE_ENCODINGS: Final = ("cl100k_base", "o200k_base", "p50k_base", "p50k_edit", "o200k_harmony") +ENCODINGS: Final = ("cl100k_base", "o200k_base", "p50k_base", "p50k_edit", "o200k_harmony") UNICODE_TEXTS: Final = ("hello world", "café 漢字 🙂", "", "a\ud800b", "\ud83d\ude42", "🙂\ud83d\ude42\udfff", " " * 64) -@pytest.mark.parametrize("name", OFFLINE_ENCODINGS) +@pytest.mark.parametrize("name", ENCODINGS) @pytest.mark.parametrize("text", UNICODE_TEXTS) def test_openai_encoding_matches_python_unicode_and_batches(name: str, text: str) -> None: - assert_openai_encoding_matches_python(name, text) - - -def assert_openai_encoding_matches_python(name: str, text: str) -> None: reference: Final = tiktoken.get_encoding(name) encoding: Final = OpenAIEncoding.from_tiktoken(name) expected: Final = reference.encode(text) @@ -309,10 +305,6 @@ def test_huggingface_batch_sequence_containers_match_python(is_pretokenized: boo @pytest.mark.parametrize("name", ("cl100k_base", "o200k_base", "p50k_edit")) def test_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None: - assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name) - - -def assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None: reference: Final = tiktoken.get_encoding(name) encoding: Final = OpenAIEncoding.from_tiktoken(name) text: Final = "hello fanta" diff --git a/tests/unit/llms/anthropic/batches/test_handler.py b/tests/unit/llms/anthropic/batches/test_handler.py index 6fde6350127..28b84123482 100644 --- a/tests/unit/llms/anthropic/batches/test_handler.py +++ b/tests/unit/llms/anthropic/batches/test_handler.py @@ -10,8 +10,7 @@ env) - and assert exactly which seam fired, with what URL/headers, and that the parsed result is the LiteLLMBatch the transform produced. The sync ``retrieve_batch`` dispatch (``_is_async`` true -> coroutine, false -> -asyncio.run) is exercised directly, mirroring the dispatch-contract discipline in -tests/test_litellm/batches/test_main.py. +asyncio.run) is exercised directly. """ from unittest.mock import AsyncMock, MagicMock, patch diff --git a/tests/unit/llms/anthropic/test_anthropic_output_format_filter.py b/tests/unit/llms/anthropic/test_anthropic_output_format_filter.py index 90cba035760..b2850192f91 100644 --- a/tests/unit/llms/anthropic/test_anthropic_output_format_filter.py +++ b/tests/unit/llms/anthropic/test_anthropic_output_format_filter.py @@ -1,9 +1,7 @@ """ Coverage for filter_anthropic_output_schema's array/object constraint stripping. -Mirrors tests/litellm/llms/anthropic/test_anthropic_schema_filter.py, but lives -under tests/test_litellm/ so the coverage-uploading CI job exercises the stripped -keyword handling (uniqueItems / contains / minProperties / maxProperties plus +Exercises the stripped keyword handling (uniqueItems / contains / minProperties / maxProperties plus multipleOf / patternProperties / propertyNames / dependentRequired / dependentSchemas / unevaluatedProperties / if / then / else / not / prefixItems), the ``uniqueItems: false`` branch, the oneOf to anyOf rewrite, and the diff --git a/tests/unit/llms/databricks/test_databricks_cost_calculator.py b/tests/unit/llms/databricks/test_databricks_cost_calculator.py index de0e547c0cd..494b99c1d11 100644 --- a/tests/unit/llms/databricks/test_databricks_cost_calculator.py +++ b/tests/unit/llms/databricks/test_databricks_cost_calculator.py @@ -16,6 +16,7 @@ NEW_MODELS: Final = ( "databricks/databricks-claude-opus-4-7", "databricks/databricks-claude-opus-4-8", "databricks/databricks-claude-opus-5", + "databricks/databricks-claude-opus-5-5", "databricks/databricks-claude-sonnet-5", "databricks/databricks-claude-fable-5", "databricks/databricks-claude-fable-5-1", @@ -32,6 +33,7 @@ PRICE_FIELDS: Final = ( "cache_read_input_token_cost", ) PUBLISHED_DBU_PER_MILLION: Final = { + "databricks/databricks-claude-opus-5-5": ("57.143", "285.714", "71.429", "2.857"), "databricks/databricks-claude-fable-5-1": ("142.858", "714.286", "178.572", "3.572"), "databricks/databricks-claude-fable-5": ("142.858", "714.286", "178.572", "14.286"), "databricks/databricks-claude-opus-5": ("71.429", "357.143", "89.286", "7.143"), @@ -118,6 +120,7 @@ def _dollars_per_token(dbu_per_million: str) -> float: [ "databricks/databricks-claude-opus-4-8", "databricks/databricks-claude-opus-5", + "databricks/databricks-claude-opus-5-5", "databricks/databricks-claude-sonnet-5", ], ) diff --git a/tests/unit/llms/langflow/chat/test_langflow_chat_transformation.py b/tests/unit/llms/langflow/chat/test_langflow_chat_transformation.py index 179a6cad4aa..138ad8bb81f 100644 --- a/tests/unit/llms/langflow/chat/test_langflow_chat_transformation.py +++ b/tests/unit/llms/langflow/chat/test_langflow_chat_transformation.py @@ -222,7 +222,8 @@ def test_langflow_extra_body_cannot_inject_tweaks_into_run_payload(): def fake_post(*args, **kwargs): body = kwargs.get("data") - posted_bodies.append(json.loads(body) if isinstance(body, str) else body) + if str(kwargs.get("url", "")).startswith("http://example.com"): + posted_bodies.append(json.loads(body) if isinstance(body, (str, bytes)) else body) resp = MagicMock(spec=httpx.Response) resp.status_code = 200 resp.json.return_value = {"outputs": [{"outputs": [{"results": {"message": {"text": "hi"}}}]}]} diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index 4f23ac1773a..0b37e033023 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -1,7 +1,12 @@ import base64 +from pathlib import Path +from typing import Final +import httpx import pytest +import respx +import litellm from litellm.litellm_core_utils.prompt_templates.factory import ( convert_to_gemini_tool_call_result, ) @@ -2727,3 +2732,40 @@ def test_gemini_server_side_tool_signature_not_duplicated_on_text(): assert "thoughtSignature" not in text_part tool_call_part = next(p for p in parts if "toolCall" in p) assert tool_call_part["thoughtSignature"] == "server_side_signature" + + +WHITE_PNG: Final = (Path(__file__).parents[4] / "white_100x100.png").read_bytes() + + +@respx.mock +def test_convert_tool_response_with_url_image(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "user_url_validation", False) + image_url: Final = "https://tool-result-images.test/gemini-tool-response.png" + respx.get(image_url).mock(return_value=httpx.Response(200, content=WHITE_PNG, headers={"content-type": "image/png"})) + tool_message: Final = { + "role": "tool", + "tool_call_id": "call_test456", + "content": [ + {"type": "text", "text": '{"url": "https://example.com"}'}, + {"type": "input_image", "image_url": image_url}, + ], + } + last_message_with_tool_calls: Final = { + "tool_calls": [ + { + "id": "call_test456", + "function": {"name": "type_text_at", "arguments": '{"x": 300, "y": 400, "text": "hello"}'}, + } + ] + } + + result: Final = convert_to_gemini_tool_call_result(tool_message, last_message_with_tool_calls) + + assert isinstance(result, list) + assert len(result) == 1 + assert "inline_data" not in result[0] + function_response: Final = result[0]["function_response"] + assert function_response["name"] == "type_text_at" + assert len(function_response["parts"]) == 1 + inline_data: Final[BlobType] = function_response["parts"][0]["inline_data"] + assert inline_data == {"data": base64.b64encode(WHITE_PNG).decode(), "mime_type": "image/png"} diff --git a/tests/test_litellm/llms/volcengine/test_volcengine_embedding.py b/tests/unit/llms/volcengine/test_volcengine_embedding.py similarity index 100% rename from tests/test_litellm/llms/volcengine/test_volcengine_embedding.py rename to tests/unit/llms/volcengine/test_volcengine_embedding.py diff --git a/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py b/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py index 6115e627b26..fb5e84c8294 100644 --- a/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py +++ b/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py @@ -3717,7 +3717,7 @@ async def test_auth_vertex_ai_route(prisma_client): @pytest.mark.asyncio -async def test_user_api_key_auth_db_unavailable(): +async def test_user_api_key_auth_db_unavailable(monkeypatch): """ Test that user_api_key_auth handles DB connection failures appropriately when: 1. DB connection fails during token validation @@ -3747,7 +3747,7 @@ async def test_user_api_key_auth_db_unavailable(): # Set up test environment setattr(litellm.proxy.proxy_server, "prisma_client", MockPrismaClient()) - setattr(litellm.proxy.proxy_server, "user_api_key_cache", MockDualCache()) + monkeypatch.setattr(litellm.proxy.proxy_server, "user_api_key_cache", MockDualCache()) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr( litellm.proxy.proxy_server, @@ -3777,7 +3777,7 @@ async def test_user_api_key_auth_db_unavailable(): @pytest.mark.asyncio -async def test_user_api_key_auth_db_unavailable_not_allowed(): +async def test_user_api_key_auth_db_unavailable_not_allowed(monkeypatch): """ Test that user_api_key_auth raises an exception when: This is default behavior @@ -3808,7 +3808,7 @@ async def test_user_api_key_auth_db_unavailable_not_allowed(): # Set up test environment setattr(litellm.proxy.proxy_server, "prisma_client", MockPrismaClient()) - setattr(litellm.proxy.proxy_server, "user_api_key_cache", MockDualCache()) + monkeypatch.setattr(litellm.proxy.proxy_server, "user_api_key_cache", MockDualCache()) setattr(litellm.proxy.proxy_server, "general_settings", {}) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") diff --git a/tests/unit/proxy/test_proxy_server.py b/tests/unit/proxy/test_proxy_server.py index eae80f311d8..8947da4d9fc 100644 --- a/tests/unit/proxy/test_proxy_server.py +++ b/tests/unit/proxy/test_proxy_server.py @@ -1,5 +1,6 @@ import os import traceback +from typing import Final from unittest import mock from dotenv import load_dotenv @@ -35,6 +36,7 @@ from fastapi import FastAPI from fastapi.testclient import TestClient from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.proxy_server import ( # Replace with the actual module where your FastAPI router is defined app, initialize, @@ -418,7 +420,7 @@ def test_chat_completion_forward_llm_provider_auth_headers( @mock_patch_acompletion() @pytest.mark.asyncio -async def test_team_disable_guardrails(mock_acompletion, client_no_auth): +async def test_team_disable_guardrails(mock_acompletion, client_no_auth, monkeypatch): """ If team not allowed to turn on/off guardrails @@ -438,8 +440,9 @@ async def test_team_disable_guardrails(mock_acompletion, client_no_auth): UserAPIKeyAuth, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth - from litellm.proxy.proxy_server import hash_token, user_api_key_cache + from litellm.proxy.proxy_server import hash_token + user_api_key_cache: Final = UserApiKeyCache() _team_id = "1234" user_key = "sk-12345678" @@ -459,7 +462,7 @@ async def test_team_disable_guardrails(mock_acompletion, client_no_auth): user_api_key_cache.set_cache(key=hash_token(user_key), value=valid_token) user_api_key_cache.set_cache(key="team_id:{}".format(_team_id), value=team_obj) - setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) + monkeypatch.setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "prisma_client", "hello-world") @@ -481,10 +484,11 @@ from tests.unit.proxy.test_custom_callback_input import CompletionCustomHandler @mock_patch_acompletion() -def test_custom_logger_failure_handler(mock_acompletion, client_no_auth): +def test_custom_logger_failure_handler(mock_acompletion, client_no_auth, monkeypatch): from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.proxy_server import hash_token, user_api_key_cache + from litellm.proxy.proxy_server import hash_token + user_api_key_cache: Final = UserApiKeyCache() rpm_limit = 0 mock_api_key = "sk-my-test-key" @@ -501,7 +505,7 @@ def test_custom_logger_failure_handler(mock_acompletion, client_no_auth): litellm.callbacks = [mock_logger, mock_logger_unit_tests] proxy_logging_obj._init_litellm_callbacks(llm_router=None) - setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) + monkeypatch.setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "prisma_client", "FAKE-VAR") setattr(litellm.proxy.proxy_server, "proxy_logging_obj", proxy_logging_obj) @@ -1296,7 +1300,7 @@ async def test_create_team_member_add(prisma_client, new_member_method): # noqa @pytest.mark.parametrize("team_route", ["/team/member_add", "/team/member_delete"]) @pytest.mark.asyncio async def test_create_team_member_add_team_admin_user_api_key_auth( - prisma_client, team_member_role, team_route # noqa: F811 # pytest fixture, not a redefinition + prisma_client, team_member_role, team_route, monkeypatch # noqa: F811 # pytest fixture, not a redefinition ): import time @@ -1307,9 +1311,10 @@ async def test_create_team_member_add_team_admin_user_api_key_auth( ProxyException, hash_token, user_api_key_auth, - user_api_key_cache, ) + user_api_key_cache: Final = UserApiKeyCache() + setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm, "max_internal_user_budget", 10) @@ -1335,7 +1340,7 @@ async def test_create_team_member_add_team_admin_user_api_key_auth( user_api_key_cache.set_cache(key="team_id:{}".format(_team_id), value=team_obj) - setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) + monkeypatch.setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) ## TEST IF TEAM ADMIN ALLOWED TO CALL /MEMBER_ADD ENDPOINT import json @@ -2349,7 +2354,7 @@ async def test_proxy_server_prisma_setup(): mock_client.db = mock_db prisma_client = await ProxyStartupEvent._setup_prisma_client( - database_url=os.getenv("DATABASE_URL"), + database_url="postgresql://user:pass@localhost:5432/litellm", proxy_logging_obj=ProxyLogging(user_api_key_cache=user_api_key_cache), user_api_key_cache=user_api_key_cache, ) @@ -2979,7 +2984,7 @@ async def test_get_config_callbacks_environment_variables(client_no_auth): @pytest.mark.asyncio -async def test_update_config_success_callback_normalization(): +async def test_update_config_success_callback_normalization(monkeypatch): """ Ensure success_callback values are normalized to lowercase when updating config. This prevents delete_callback (which searches lowercase) from failing on mixed case inputs like 'SQS'. @@ -2987,7 +2992,7 @@ async def test_update_config_success_callback_normalization(): import litellm.proxy.proxy_server as proxy_server from litellm.proxy._types import ConfigYAML - setattr(proxy_server, "proxy_logging_obj", MagicMock()) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock()) existing_litellm_settings = {"success_callback": ["langfuse"]} @@ -3013,7 +3018,7 @@ async def test_update_config_success_callback_normalization(): self.db.litellm_config.find_first = AsyncMock(side_effect=fake_find_first) self.db.litellm_config.upsert = AsyncMock(side_effect=fake_upsert) - setattr(proxy_server, "prisma_client", MockPrisma()) + monkeypatch.setattr(proxy_server, "prisma_client", MockPrisma()) class MockProxyConfig: async def add_deployment(self, prisma_client=None, proxy_logging_obj=None): # noqa: F811 # pytest fixture, not a redefinition @@ -3022,7 +3027,7 @@ async def test_update_config_success_callback_normalization(): def reject_config_owned_writes(self, *, section_name, changed_keys): return None - setattr(proxy_server, "proxy_config", MockProxyConfig()) + monkeypatch.setattr(proxy_server, "proxy_config", MockProxyConfig()) config_update = ConfigYAML(litellm_settings={"success_callback": ["SQS", "sQs"]}) from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth diff --git a/tests/unit/router_strategy/test_complexity_router.py b/tests/unit/router_strategy/test_complexity_router.py index 3401a335b2f..2d66524326f 100644 --- a/tests/unit/router_strategy/test_complexity_router.py +++ b/tests/unit/router_strategy/test_complexity_router.py @@ -14120,6 +14120,7 @@ class TestContextWindowEscalation: assert result.routing_decision["context_escalated"] is True @pytest.mark.asyncio + @pytest.mark.usefixtures("local_model_cost_map") @pytest.mark.parametrize( "deployments,tiers,expected_model", [ diff --git a/tests/unit/rust_bridge/test_callbacks_legacy_python.py b/tests/unit/rust_bridge/test_callbacks_legacy_python.py index 05f2d13a079..7365679a28c 100644 --- a/tests/unit/rust_bridge/test_callbacks_legacy_python.py +++ b/tests/unit/rust_bridge/test_callbacks_legacy_python.py @@ -8,11 +8,10 @@ from typing import Final import pytest from pydantic import TypeAdapter -import litellm from litellm._internal_context import is_internal_call from litellm.litellm_core_utils.litellm_logging import Logging from litellm.rust_bridge import callbacks_legacy_python as legacy -from litellm.rust_bridge.callbacks_legacy_python import check_limits, failure_handler, setup +from litellm.rust_bridge.callbacks_legacy_python import failure_handler, setup _OCR_KWARGS: Final = MappingProxyType( { @@ -22,33 +21,6 @@ _OCR_KWARGS: Final = MappingProxyType( ) -@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"]) -@pytest.mark.parametrize( - "cap, request_retry_count, refused", - [(5, 5, True), (5, 4, False), (0, 0, False), (0, 1, True)], - ids=[ - "cap-above-four-reached", - "cap-above-four-not-reached", - "first-attempt-passes-cap-of-zero", - "cap-of-zero-refuses-first-retry", - ], -) -def test_check_limits_reads_request_retry_count( - monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, request_retry_count: int, refused: bool -) -> None: - monkeypatch.setattr(litellm, "num_retries_per_request", cap) - monkeypatch.setattr(litellm, "max_budget", None) - kwargs: Final = { - "model": "mistral/mistral-ocr-latest", - metadata_key: {"request_retry_count": request_retry_count}, - } - if refused: - with pytest.raises(RuntimeError, match="Max retries per request hit!"): - check_limits(kwargs) - else: - check_limits(kwargs) - - def _supplied_logger() -> Logging: return Logging( model="mistral/mistral-ocr-latest", diff --git a/tests/unit/rust_bridge/test_catalog.py b/tests/unit/rust_bridge/test_catalog.py index eeea364b674..2b3edac612e 100644 --- a/tests/unit/rust_bridge/test_catalog.py +++ b/tests/unit/rust_bridge/test_catalog.py @@ -50,12 +50,13 @@ def test_shipped_decisions( monkeypatch.setenv("LITELLM_RUST", environment) context: Final = RouteContext(route, provider=provider, model="test-model", delivery=delivery) - if route is Route.OCR: - assert catalog.rollout(context) is Rollout.RUST_REQUIRED - assert catalog.decision(context) is Decision.RUST_REQUIRED - elif route is Route.TRANSCRIPTION and provider == "bedrock": + if route is Route.OCR or (route is Route.TRANSCRIPTION and provider == "bedrock"): assert catalog.rollout(context) is Rollout.RUST_REQUIRED assert catalog.decision(context) is Decision.RUST_REQUIRED + elif route is Route.MESSAGES and provider == "anthropic": + assert catalog.rollout(context) is Rollout.RUST_OPT_IN + opted_in: Final = environment == "1" or (environment is None and process is True) + assert catalog.decision(context) is (Decision.RUST_WITH_FALLBACK if opted_in else Decision.PYTHON) else: assert catalog.rollout(context) is Rollout.PYTHON_ONLY assert catalog.decision(context) is Decision.PYTHON @@ -145,10 +146,7 @@ def test_ocr_has_no_python_path_to_opt_out_to( monkeypatch.setenv("LITELLM_RUST", environment) assert catalog.decision(RouteContext(Route.OCR, model="m")) is Decision.RUST_REQUIRED - assert ( - catalog.decision(RouteContext(Route.OCR, provider="aws_textract", model="m")) - is Decision.RUST_REQUIRED - ) + assert catalog.decision(RouteContext(Route.OCR, provider="aws_textract", model="m")) is Decision.RUST_REQUIRED @pytest.mark.parametrize( diff --git a/tests/unit/rust_bridge/test_preflight.py b/tests/unit/rust_bridge/test_preflight.py new file mode 100644 index 00000000000..a8b1a40a00b --- /dev/null +++ b/tests/unit/rust_bridge/test_preflight.py @@ -0,0 +1,63 @@ +import inspect +from collections.abc import Callable +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import pytest +from pydantic import TypeAdapter + +import litellm +from litellm.rust_bridge import preflight +from litellm.rust_bridge.preflight import check_limits + + +@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"]) +@pytest.mark.parametrize( + "cap, request_retry_count, refused", + [(5, 5, True), (5, 4, False), (0, 0, False), (0, 1, True)], + ids=[ + "cap-above-four-reached", + "cap-above-four-not-reached", + "first-attempt-passes-cap-of-zero", + "cap-of-zero-refuses-first-retry", + ], +) +def test_check_limits_reads_request_retry_count( + monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, request_retry_count: int, refused: bool +) -> None: + monkeypatch.setattr(litellm, "num_retries_per_request", cap) + monkeypatch.setattr(litellm, "max_budget", None) + kwargs: Final = { + "model": "mistral/mistral-ocr-latest", + metadata_key: {"request_retry_count": request_retry_count}, + } + if refused: + with pytest.raises(RuntimeError, match="Max retries per request hit!"): + check_limits(kwargs) + else: + check_limits(kwargs) + + +def test_check_limits_refuses_a_call_over_the_budget(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "num_retries_per_request", None) + monkeypatch.setattr(litellm, "max_budget", 1.0) + monkeypatch.setattr(litellm, "_current_cost", 1.5) + with pytest.raises(litellm.BudgetExceededError): + check_limits({"model": "mistral/mistral-ocr-latest"}) + + +CONTRACT_PATH: Final = Path(__file__).parents[3] / "litellm-rust/crates/python-bridge/preflight_contract.json" +_SHIMS: Final[MappingProxyType[str, Callable[..., object]]] = MappingProxyType( + { + "credential_list": preflight.credential_list, + "warn_unknown_credential": preflight.warn_unknown_credential, + "check_limits": preflight.check_limits, + } +) + + +def test_the_rust_contract_matches_the_shim_signatures() -> None: + contract: Final = TypeAdapter(dict[str, list[str]]).validate_json(CONTRACT_PATH.read_text()) + + assert contract == {name: list(inspect.signature(_SHIMS[name]).parameters) for name in contract} diff --git a/tests/unit/secret_managers/test_cyberark_secret_manager.py b/tests/unit/secret_managers/test_cyberark_secret_manager.py index 3f3669ab9ef..334e6437f1d 100644 --- a/tests/unit/secret_managers/test_cyberark_secret_manager.py +++ b/tests/unit/secret_managers/test_cyberark_secret_manager.py @@ -1,7 +1,9 @@ +import asyncio import json from pathlib import Path from typing import Final, TypedDict, cast +import httpx import pytest import respx @@ -107,6 +109,86 @@ async def test_async_write_matches_parity_fixture(monkeypatch: pytest.MonkeyPatc assert value_route.calls.last.request.content == b"v" +@pytest.mark.asyncio +@respx.mock +async def test_async_write_retries_policy_load_conflict(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + fixture: Final = _fixture() + manager: Final = _configure_manager(monkeypatch, fixture) + secret: Final = fixture["secrets"][0] + endpoint: Final = fixture["endpoint"] + _respond(respx.post(endpoint + fixture["authenticate_path"]), content=fixture["token_json"].encode()) + policy_route: Final = respx.post(endpoint + fixture["policy_path"]).mock( + side_effect=[httpx.Response(409), httpx.Response(409), httpx.Response(201)] + ) + value_route: Final = respx.post(endpoint + secret["path"]).mock( + side_effect=lambda _: httpx.Response(201 if policy_route.call_count == 3 else 404) + ) + + result: Final = await manager.async_write_secret(secret["name"], "v") # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # legacy secret manager API is untyped + + assert policy_route.call_count == 3 + assert value_route.call_count == 1 + assert result["status"] == "success" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "policy_outcome", + [422, 500, httpx.ConnectError("conjur unreachable")], + ids=["unprocessable", "server_error", "unreachable"], +) +@respx.mock +async def test_async_write_does_not_retry_non_conflict_policy_failures( + monkeypatch: pytest.MonkeyPatch, policy_outcome: int | httpx.ConnectError +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + fixture: Final = _fixture() + manager: Final = _configure_manager(monkeypatch, fixture) + secret: Final = fixture["secrets"][0] + endpoint: Final = fixture["endpoint"] + _respond(respx.post(endpoint + fixture["authenticate_path"]), content=fixture["token_json"].encode()) + policy_route: Final = respx.post(endpoint + fixture["policy_path"]) + if isinstance(policy_outcome, int): + _respond(policy_route, status_code=policy_outcome) + else: + policy_route.mock(side_effect=policy_outcome) + value_route: Final = _respond(respx.post(endpoint + secret["path"]), status_code=201) + + await manager.async_write_secret(secret["name"], "v") # pyright: ignore[reportUnknownMemberType] # legacy secret manager API is untyped + + assert policy_route.call_count == 1 + assert value_route.call_count == 1 + + +@pytest.mark.asyncio +@respx.mock +async def test_concurrent_async_writes_load_policy_one_at_a_time(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + fixture: Final = _fixture() + manager: Final = _configure_manager(monkeypatch, fixture) + endpoint: Final = fixture["endpoint"] + _respond(respx.post(endpoint + fixture["authenticate_path"]), content=fixture["token_json"].encode()) + in_flight: Final = asyncio.Semaphore(1) + + async def load_policy(_: httpx.Request) -> httpx.Response: + if in_flight.locked(): + return httpx.Response(409) + async with in_flight: + await asyncio.sleep(0.05) + return httpx.Response(201) + + policy_route: Final = respx.post(endpoint + fixture["policy_path"]).mock(side_effect=load_policy) + respx.post(url__startswith=endpoint + "/secrets/").respond(status_code=201) # pyright: ignore[reportUnknownMemberType] # respx route stubs leave response builder partially unknown + + results: Final = await asyncio.gather( + *(manager.async_write_secret(f"concurrent-{index}", "v") for index in range(4)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # legacy secret manager API is untyped + ) + + assert policy_route.call_count == 4 + assert [result["status"] for result in results] == ["success"] * 4 + + def test_missing_credentials_raise_value_error(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True) for name in ( diff --git a/tests/unit/test_logging.py b/tests/unit/test_logging.py index d9cfe88d52f..5cdecc31575 100644 --- a/tests/unit/test_logging.py +++ b/tests/unit/test_logging.py @@ -9,7 +9,7 @@ import sys import time from io import StringIO from pathlib import Path -from typing import List +from typing import Final, List import pytest from pydantic import BaseModel, computed_field @@ -1584,6 +1584,8 @@ def _emit_access_line(full_path: str) -> str: handler = logging.StreamHandler(stream) handler.setFormatter(AccessFormatter('%(client_addr)s - "%(request_line)s" %(status_code)s', use_colors=False)) saved_level, saved_propagate = logger.level, logger.propagate + saved_filters: Final = logger.filters[:] + logger.filters = [f for f in saved_filters if type(f).__module__.split(".")[0] == "litellm"] logger.addHandler(handler) logger.setLevel(logging.INFO) logger.propagate = False @@ -1593,6 +1595,7 @@ def _emit_access_line(full_path: str) -> str: logger.removeHandler(handler) logger.setLevel(saved_level) logger.propagate = saved_propagate + logger.filters = saved_filters return stream.getvalue() diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py index c06216e4f4e..57200a79a8c 100644 --- a/tests/unit/test_main.py +++ b/tests/unit/test_main.py @@ -17,6 +17,7 @@ import respx import urllib.parse from importlib import import_module +from pathlib import Path from unittest.mock import MagicMock, patch import litellm @@ -56,6 +57,9 @@ def add_api_keys_to_env(monkeypatch): monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False) +WHITE_PNG: Final = (Path(__file__).parents[1] / "white_100x100.png").read_bytes() + + @pytest.fixture def openai_api_response(): mock_response_data = { @@ -213,6 +217,102 @@ async def test_url_with_format_param_openai(model, sync_mode): assert "format" not in json_str +@pytest.mark.parametrize( + "model", + [ + "gemini/gemini-1.5-flash", + "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + "bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0", + "anthropic/claude-3-5-sonnet", + ], +) +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_url_with_format_param(model, sync_mode, monkeypatch): + from litellm import acompletion, completion + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + + if sync_mode: + client = HTTPHandler() + else: + client = AsyncHTTPHandler() + + image_url: Final = ( + "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png" + f"?case={sync_mode}-{model}" + ) + args = { + "model": model, + "messages": [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": { + "url": image_url, + "format": "image/png", + }, + }, + {"type": "text", "text": "Describe this image"}, + ], + } + ], + } + if model.startswith("gemini/"): + args["api_key"] = "test-api-key" + monkeypatch.setattr(litellm, "user_url_validation", False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "module_level_aclient", AsyncHTTPHandler(transport=httpx.AsyncHTTPTransport())) + with ( + respx.mock(assert_all_called=False) as image_host, + patch.object(client, "post", new=MagicMock()) as mock_client, + ): + image_route = image_host.get(image_url).mock( + return_value=httpx.Response(200, content=WHITE_PNG, headers={"content-type": "image/png"}) + ) + try: + if sync_mode: + response = completion(**args, client=client) + else: + response = await acompletion(**args, client=client) + print(response) + except Exception as e: + pass + + mock_client.assert_called() + + print(mock_client.call_args.kwargs) + + if "data" in mock_client.call_args.kwargs: + json_str = mock_client.call_args.kwargs["data"] + else: + json_str = json.dumps(mock_client.call_args.kwargs["json"]) + + if isinstance(json_str, bytes): + json_str = json_str.decode("utf-8") + + print(f"type of json_str: {type(json_str)}") + + if model.startswith("bedrock/invoke/"): + assert "https://awsmp-logos.s3.amazonaws.com" not in json_str + assert '"type":"base64"' in json_str or '"type": "base64"' in json_str + assert '"data"' in json_str + elif model.startswith("bedrock/"): + assert "https://awsmp-logos.s3.amazonaws.com" not in json_str + assert '"bytes"' in json_str or '"bytes":' in json_str + elif model.startswith("anthropic/"): + assert "https://awsmp-logos.s3.amazonaws.com" in json_str + assert '"type":"url"' in json_str or '"type": "url"' in json_str + else: + assert "png" in json_str + assert "jpeg" not in json_str + + fetches_image: Final = not model.startswith("anthropic/") + assert image_route.called is fetches_image + assert (base64.b64encode(WHITE_PNG).decode() in json_str) is fetches_image + + def test_bedrock_latency_optimized_inference(): from litellm.llms.custom_httpx.http_handler import HTTPHandler diff --git a/tests/unit/test_model_prices_schema.py b/tests/unit/test_model_prices_schema.py index 052278631e2..a05c345b5cd 100644 --- a/tests/unit/test_model_prices_schema.py +++ b/tests/unit/test_model_prices_schema.py @@ -266,6 +266,61 @@ def test_openai_reasoning_family_entries_carry_supports_reasoning(prices: dict): ) +_ABSENT: Final = object() + +REASONING_ANNOTATION_KEYS: Final = ( + "supports_reasoning", + "supports_minimal_reasoning_effort", + "supports_none_reasoning_effort", + "supports_xhigh_reasoning_effort", + "default_reasoning_effort", +) + + +def chatgpt_openai_twins(prices: dict) -> list[tuple[str, str]]: + """`chatgpt/` rows paired with the bare `` row served by the openai provider. + + Scoped to openai twins on purpose. `ChatGPTConfig` and `ChatGPTResponsesAPIConfig` subclass + their openai counterparts, so a chatgpt row's reasoning behaviour is whatever the openai row + describes. The azure rows are a separate registry that already diverges from openai here, and + pinning them to each other would assert something this repository does not control. + """ + pairs = [] + for name, entry in prices.items(): + if not isinstance(entry, dict) or not name.startswith("chatgpt/"): + continue + bare = name.split("/", 1)[1] + twin = prices.get(bare) + if isinstance(twin, dict) and twin.get("litellm_provider") == "openai": + pairs.append((name, bare)) + return pairs + + +def test_chatgpt_rows_carry_their_openai_twin_reasoning_annotations(prices: dict): + """A chatgpt row must not silently drop the reasoning annotations of the model it proxies. + + `litellm.utils._get_model_info_from_generalization` refuses to fall back when an exact cost-map + key exists, so an unannotated `chatgpt/` row wins over its annotated twin and + `/model/info` reports the model as non-reasoning. + """ + twins = chatgpt_openai_twins(prices) + assert twins, "no chatgpt/* row has an openai twin any more; this guard has stopped guarding" + + mismatched = [] + for name, bare in twins: + for key in REASONING_ANNOTATION_KEYS: + if prices[name].get(key, _ABSENT) != prices[bare].get(key, _ABSENT): + mismatched.append( + f"{name}.{key} is {prices[name].get(key)!r}, {bare}.{key} is {prices[bare].get(key)!r}" + ) + + assert mismatched == [], ( + "chatgpt/* entries proxy their openai twin through ChatGPTConfig, so they must carry the " + "same reasoning annotations; an exact cost-map key blocks the generalization fallback, so " + "a missing flag here is reported to callers as 'not a reasoning model':\n" + "\n".join(mismatched) + ) + + def test_chat_latest_declares_the_one_effort_openai_accepts(prices: dict): """OpenAI rejects every reasoning.effort on chat-latest except medium, and a reasoning entry with no declared levels resolves to None, which lets /model_group/info and the dashboard offer diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 3393c2f0d3c..e99cefb35dd 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -8,7 +8,7 @@ import os import sys import threading import warnings -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import AsyncIterable, AsyncIterator, Awaitable, Callable, Mapping from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final, Literal @@ -37,6 +37,7 @@ from litellm.models.access_group import LiteLLM_AccessGroupTable from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, Member, ProxyException, UserAPIKeyAuth from litellm.router import ( MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS, + MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS, FallbackAwareAnthropicMessagesStream, _anthropic_stream_commits_now, _anthropic_stream_error_is_gateway_verdict, @@ -46,6 +47,7 @@ from litellm.router import ( _anthropic_stream_should_decline_fallback, _anthropic_stream_should_drop_pre_content_ping, _is_retriable_anthropic_status, + _responses_stream_holds_event, ) from litellm.router_strategy import simple_shuffle from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit @@ -4213,6 +4215,7 @@ async def test_aresponses_streaming_iterator_fallback(): hidden_params={"model_id": "src-deployment-1"}, ) fallback_chunks = [ + MagicMock(type="response.created"), MagicMock(type="response.output_text.delta"), MagicMock(type="response.completed"), ] @@ -4235,7 +4238,7 @@ async def test_aresponses_streaming_iterator_fallback(): assert wrapped._hidden_params.get("model_id") == "src-deployment-1" collected = [c async for c in wrapped] - assert len(collected) == 3 # 1 primary chunk + 2 fallback chunks + assert collected == fallback_chunks call_kwargs = mock_fallback_utils.call_args.kwargs fbk = call_kwargs["kwargs"] # Bound methods compare equal when they share the same instance + __func__. @@ -4522,20 +4525,35 @@ def _make_native_responses_iterator(*, sse_payloads: tuple[dict[str, str], ...], _RESPONSES_LIFECYCLE_PAYLOADS: Final = ({"type": "response.created"}, {"type": "response.in_progress"}) +async def _events_until_error(stream: AsyncIterable[object]) -> AsyncIterator[object]: + try: + async for chunk in stream: + yield chunk + except Exception as error: + yield error + + @pytest.mark.asyncio async def test_aresponses_streaming_iterator_falls_back_on_transport_drop_before_output(): """A connection lost after response.created but before any output item is re-routed to the - fallback with the original input, the same as a provider error event would be.""" + fallback with the original input, the same as a provider error event would be, and the client + sees one response lifecycle: the fallback's, whose id the completed event carries.""" router: Final = _make_router_with_fallback() src: Final = _make_native_responses_iterator( sse_payloads=_RESPONSES_LIFECYCLE_PAYLOADS, trailing_error=httpx.ReadError("Response payload is not completed"), ) + fallback_chunks: Final = [ + MagicMock(type="response.created", response=MagicMock(id="resp_fallback")), + MagicMock(type="response.in_progress", response=MagicMock(id="resp_fallback")), + MagicMock(type="response.output_text.delta"), + MagicMock(type="response.completed", response=MagicMock(id="resp_fallback")), + ] with patch.object( router, "async_function_with_fallbacks_common_utils", - return_value=_AsyncList([MagicMock(type="response.completed")]), + return_value=_AsyncList(fallback_chunks), ) as mock_fallback_utils: wrapped: Final = await router._aresponses_streaming_iterator( response=src, @@ -4546,9 +4564,10 @@ async def test_aresponses_streaming_iterator_falls_back_on_transport_drop_before "original_generic_function": litellm.aresponses, }, ) - seen: Final = [chunk.type async for chunk in wrapped] + collected: Final = [chunk async for chunk in wrapped] - assert seen == ["response.created", "response.in_progress", "response.completed"] + assert collected == fallback_chunks + assert [chunk.response.id for chunk in collected if chunk.type == "response.created"] == ["resp_fallback"] assert isinstance(mock_fallback_utils.call_args.kwargs["e"], MidStreamFallbackError) assert mock_fallback_utils.call_args.kwargs["kwargs"]["input"] == "Hello" @@ -4576,17 +4595,197 @@ async def test_aresponses_streaming_iterator_surfaces_transport_drop_when_no_fal "original_generic_function": litellm.aresponses, }, ) - with pytest.raises(httpx.ReadError) as exc_info: - async for _ in wrapped: - pass + outcome: Final = [item async for item in _events_until_error(wrapped)] - assert exc_info.value is transport_error + assert [item.type for item in outcome[:-1]] == ["response.created", "response.in_progress"] + assert outcome[-1] is transport_error assert mock_fallback_utils.await_count == 1 trigger: Final = mock_fallback_utils.await_args.kwargs["e"] assert isinstance(trigger, MidStreamFallbackError) assert trigger.original_exception is transport_error +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_forwards_lifecycle_events_in_order_once_output_starts(): + router: Final = _make_router_with_fallback() + chunks: Final = [ + MagicMock(type="response.created"), + MagicMock(type="response.in_progress"), + MagicMock(type="response.output_text.delta"), + MagicMock(type="response.completed"), + ] + wrapped: Final = await router._aresponses_streaming_iterator( + response=_make_responses_iterator(chunks=chunks), + initial_kwargs={"model": "gpt-4", "stream": True, "input": "Hello"}, + ) + + assert [chunk async for chunk in wrapped] == chunks + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_flushes_held_lifecycle_events_when_the_stream_ends_without_output(): + router: Final = _make_router_with_fallback() + chunks: Final = [MagicMock(type="response.created"), MagicMock(type="response.in_progress")] + wrapped: Final = await router._aresponses_streaming_iterator( + response=_make_responses_iterator(chunks=chunks), + initial_kwargs={"model": "gpt-4", "stream": True, "input": "Hello"}, + ) + + assert [chunk async for chunk in wrapped] == chunks + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_forwards_held_lifecycle_events_before_a_non_fallback_error(): + router: Final = _make_router_with_fallback() + chunks: Final = [MagicMock(type="response.created"), MagicMock(type="response.in_progress")] + client_error: Final = litellm.BadRequestError(message="bad input", model="gpt-4", llm_provider="openai") + wrapped: Final = await router._aresponses_streaming_iterator( + response=_make_responses_iterator(chunks=chunks, error=client_error), + initial_kwargs={"model": "gpt-4", "stream": True, "input": "Hello"}, + ) + outcome: Final = [item async for item in _events_until_error(wrapped)] + + assert outcome[:-1] == chunks + assert outcome[-1] is client_error + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_commits_held_lifecycle_events_at_the_hold_cap(): + router: Final = _make_router_with_fallback() + chunks: Final = [MagicMock(type="response.in_progress") for _ in range(MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS + 1)] + src: Final = _make_responses_iterator( + chunks=chunks, + error=MidStreamFallbackError( + message="dropped before output", model="gpt-4", llm_provider="openai", is_pre_first_chunk=True + ), + ) + fallback_chunks: Final = [MagicMock(type="response.created"), MagicMock(type="response.completed")] + + with patch.object(router, "async_function_with_fallbacks_common_utils", return_value=_AsyncList(fallback_chunks)): + wrapped: Final = await router._aresponses_streaming_iterator( + response=src, + initial_kwargs={"model": "gpt-4", "stream": True, "input": "Hello"}, + ) + collected: Final = [chunk async for chunk in wrapped] + + assert collected == [*chunks, *fallback_chunks] + + +@pytest.mark.parametrize( + ("event_type", "held_event_count", "expected"), + [ + ("response.created", 0, True), + ("response.in_progress", 1, True), + ("response.queued", MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS - 1, True), + ("response.in_progress", MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS, False), + ("response.output_item.added", 0, False), + ("response.output_text.delta", 0, False), + ("response.completed", 0, False), + ], +) +def test_responses_stream_holds_event_holds_only_pre_output_lifecycle_events_under_the_cap( + event_type: str, held_event_count: int, expected: bool +): + assert _responses_stream_holds_event(MagicMock(type=event_type), held_event_count) is expected + + +@pytest.mark.asyncio +async def test_aresponses_fallback_attempt_drops_held_lifecycle_events_when_a_fallback_lands(): + router: Final = _make_router_with_fallback() + trigger: Final = MidStreamFallbackError( + message="dropped before output", model="gpt-4", llm_provider="openai", is_pre_first_chunk=True + ) + held: Final = (MagicMock(type="response.created"), MagicMock(type="response.in_progress")) + fallback_chunks: Final = [MagicMock(type="response.created"), MagicMock(type="response.completed")] + adopt_headers: Final = MagicMock(return_value=({}, {})) + + with patch.object( + router, "async_function_with_fallbacks_common_utils", return_value=_AsyncList(fallback_chunks) + ) as mock_fallback_utils: + collected: Final = [ + chunk + async for chunk in router._aresponses_fallback_attempt( + trigger, + _make_responses_iterator(), + {"model": "gpt-4", "stream": True, "input": "Hello"}, + adopt_headers, + held, + ) + ] + + assert collected == fallback_chunks + adopt_headers.assert_called_once() + assert mock_fallback_utils.await_args.kwargs["e"] is trigger + + +@pytest.mark.asyncio +async def test_aresponses_fallback_attempt_replays_held_lifecycle_events_when_the_fallback_dies_before_its_first_event(): + """A fallback stream that raises before yielding anything announced no response of its own, so the + primary's held created/in_progress pair is replayed ahead of the error and the client sees the + announcement the failure belongs to, the same as when no fallback was attempted at all.""" + router: Final = _make_router_with_fallback() + trigger: Final = MidStreamFallbackError( + message="dropped before output", model="gpt-4", llm_provider="openai", is_pre_first_chunk=True + ) + held: Final = (MagicMock(type="response.created"), MagicMock(type="response.in_progress")) + fallback_error: Final = RuntimeError("fallback closed before its first event") + adopt_headers: Final = MagicMock(return_value=({}, {})) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=_make_responses_iterator(error=fallback_error), + ): + outcome: Final = [ + item + async for item in _events_until_error( + router._aresponses_fallback_attempt( + trigger, + _make_responses_iterator(), + {"model": "gpt-4", "stream": True, "input": "Hello"}, + adopt_headers, + held, + ) + ) + ] + + assert outcome == [*held, fallback_error] + + +@pytest.mark.asyncio +async def test_aresponses_fallback_attempt_does_not_replay_held_lifecycle_events_once_the_fallback_announced_itself(): + """Once the fallback has yielded its own created event, a later failure must not replay the + primary's held pair on top of it, or the client would again see two announced response ids.""" + router: Final = _make_router_with_fallback() + trigger: Final = MidStreamFallbackError( + message="dropped before output", model="gpt-4", llm_provider="openai", is_pre_first_chunk=True + ) + held: Final = (MagicMock(type="response.created"), MagicMock(type="response.in_progress")) + fallback_created: Final = MagicMock(type="response.created") + fallback_error: Final = RuntimeError("fallback dropped after announcing itself") + adopt_headers: Final = MagicMock(return_value=({}, {})) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=_make_responses_iterator(chunks=(fallback_created,), error=fallback_error), + ): + outcome: Final = [ + item + async for item in _events_until_error( + router._aresponses_fallback_attempt( + trigger, + _make_responses_iterator(), + {"model": "gpt-4", "stream": True, "input": "Hello"}, + adopt_headers, + held, + ) + ) + ] + + assert outcome == [fallback_created, fallback_error] + + @pytest.mark.asyncio async def test_aresponses_streaming_iterator_partial_content_injects_continuation(): """Mid-stream error: input is rewritten to include user prompt + diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 0cdc52c9a93..5b327305e31 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -766,6 +766,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "cache_creation_input_token_cost_above_272k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_272k_tokens_flex": {"type": "number"}, "cache_creation_input_token_cost_above_272k_tokens_priority": {"type": "number"}, + "cache_creation_input_token_cost_above_200k_tokens_batches": {"type": "number"}, "cache_creation_input_token_cost_above_272k_tokens_batches": {"type": "number"}, "cache_creation_input_token_cost_batches": {"type": "number"}, "cache_creation_input_token_cost_flex": {"type": "number"}, diff --git a/tests/unit/test_video_generation.py b/tests/unit/test_video_generation.py index 644c7a41f49..5c1d0bfa884 100644 --- a/tests/unit/test_video_generation.py +++ b/tests/unit/test_video_generation.py @@ -11,6 +11,7 @@ import litellm from litellm.cost_calculator import default_video_cost_calculator from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.llms.gemini.videos.transformation import GeminiVideoConfig @@ -988,6 +989,7 @@ class TestVideoLogging: """ custom_logger = self.TestVideoLogger() litellm.logging_callback_manager._reset_all_callbacks() + await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0) litellm.callbacks = [custom_logger] # Mock video generation response diff --git a/tests/white_100x100.png b/tests/white_100x100.png new file mode 100644 index 00000000000..fdd268ded88 Binary files /dev/null and b/tests/white_100x100.png differ diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx index 37f63433d2a..a693ee971d4 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx @@ -2338,6 +2338,103 @@ describe("TeamInfoView - the exact bytes the update call sends", () => { expect(wireBody(payload)).toStrictEqual(expected); }); + const openEditorWithMemberBudgetAlerts = async (user: ReturnType) => { + vi.mocked(networking.teamInfoCall).mockResolvedValue( + createMockTeamData({ + models: ["gpt-4"], + team_member_budget_table: { max_budget: 42 }, + metadata: { team_member_max_budget_alert_emails: { "50": [], "100": ["finance@test.com"] } }, + }), + ); + vi.mocked(networking.teamUpdateCall).mockResolvedValue({ data: {}, team_id: "123" } as any); + + renderWithProviders(); + await waitFor(() => expect(screen.queryAllByText("Test Team").length).toBeGreaterThan(0)); + await user.click(screen.getByRole("tab", { name: "Settings" })); + await user.click(await screen.findByRole("button", { name: /edit settings/i })); + await screen.findByLabelText("Team Name"); + }; + + const memberBudgetAlertEmails = (payload: Record) => + (wireBody(payload).metadata as Record).team_member_max_budget_alert_emails; + + it("resends the stored team member budget alert thresholds when the section stays closed", async () => { + const user = userEvent.setup({ delay: null }); + await openEditorWithMemberBudgetAlerts(user); + + const payload = await save(user); + + expect(memberBudgetAlertEmails(payload)).toStrictEqual({ "50": [], "100": ["finance@test.com"] }); + }); + + it("sends the edited team member budget alert thresholds as a percent to recipients map", async () => { + const user = userEvent.setup({ delay: null }); + await openEditorWithMemberBudgetAlerts(user); + + await user.click(screen.getByText("Team Member Settings")); + await screen.findByLabelText("Default Budget (USD)"); + const thresholds = screen.getAllByPlaceholderText("% of budget"); + const recipients = screen.getAllByPlaceholderText(/Additional recipients/); + expect(thresholds.map((input) => (input as HTMLInputElement).value)).toStrictEqual(["50", "100"]); + expect(recipients.map((input) => (input as HTMLInputElement).value)).toStrictEqual(["", "finance@test.com"]); + + fireEvent.change(thresholds[0], { target: { value: "75" } }); + fireEvent.change(recipients[0], { target: { value: " lead@test.com, finance@test.com " } }); + await user.click(screen.getByRole("button", { name: "Add Budget Alert Threshold" })); + fireEvent.change(screen.getAllByPlaceholderText("% of budget")[2], { target: { value: "90" } }); + + const payload = await save(user); + + expect(memberBudgetAlertEmails(payload)).toStrictEqual({ + "75": ["lead@test.com", "finance@test.com"], + "100": ["finance@test.com"], + "90": [], + }); + }); + + it("drops the team member budget alert thresholds key once every row is removed", async () => { + const user = userEvent.setup({ delay: null }); + await openEditorWithMemberBudgetAlerts(user); + + await user.click(screen.getByText("Team Member Settings")); + await screen.findByLabelText("Default Budget (USD)"); + const removeButtons = screen.getAllByRole("button", { name: "Remove budget alert threshold" }); + await user.click(removeButtons[1]); + await user.click(removeButtons[0]); + + const payload = await save(user); + + expect(memberBudgetAlertEmails(payload)).toBeUndefined(); + }); + + it("blocks the save when a team member budget alert threshold is above 100", async () => { + const user = userEvent.setup({ delay: null }); + await openEditorWithMemberBudgetAlerts(user); + + await user.click(screen.getByText("Team Member Settings")); + await screen.findByLabelText("Default Budget (USD)"); + const threshold = screen.getAllByPlaceholderText("% of budget")[0] as HTMLInputElement; + fireEvent.change(threshold, { target: { value: "150" } }); + expect(threshold.validity.rangeOverflow).toBe(true); + + await user.click(screen.getByRole("button", { name: /save changes/i })); + + await waitFor(() => expect(networking.teamUpdateCall).not.toHaveBeenCalled()); + }); + + it("refuses to save a team member budget alert row with no threshold", async () => { + const user = userEvent.setup({ delay: null }); + await openEditorWithMemberBudgetAlerts(user); + + await user.click(screen.getByText("Team Member Settings")); + await screen.findByLabelText("Default Budget (USD)"); + await user.click(screen.getByRole("button", { name: "Add Budget Alert Threshold" })); + await user.click(screen.getByRole("button", { name: /save changes/i })); + + await screen.findByText("Enter a whole number from 1 to 100"); + expect(networking.teamUpdateCall).not.toHaveBeenCalled(); + }); + it("carries every typed value to the update payload at the type and shape antd sends today", async () => { const user = userEvent.setup({ delay: null }); await openEditor(user); diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index 78ed507216a..3845f94593d 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -118,6 +118,13 @@ import { TEAM_INFO_TAB_LABELS, } from "./tabVisibilityUtils"; import TeamMembersComponent from "./TeamMemberTab"; +import { + isValidThreshold, + TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY, + teamMemberBudgetAlertEmailsFromRows, + teamMemberBudgetAlertRowsFromMetadata, + teamMemberBudgetAlertSummary, +} from "./teamMemberBudgetAlertEmails"; import { TeamVirtualKeysTable } from "./TeamVirtualKeysTable"; import ResetMemberBudgetsDialog from "./ResetMemberBudgetsDialog"; import { customBudgetMemberUserIds, shouldPromptMemberBudgetReset } from "./memberBudgetReset"; @@ -128,6 +135,7 @@ const UI_MANAGED_METADATA_KEYS: ReadonlySet = new Set([ "logging", "secret_manager_settings", "soft_budget_alerting_emails", + TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY, "model_tpm_limit", "model_rpm_limit", "default_estimated_output_tokens", @@ -355,6 +363,18 @@ const teamUpdateFieldsSchema = z.object({ team_member_key_duration: z.string().optional(), team_member_tpm_limit: numericInputSchema, team_member_rpm_limit: numericInputSchema, + team_member_max_budget_alert_emails: z + .array(z.object({ threshold: z.number().nullable(), emails: z.string() })) + .superRefine((rows, ctx) => { + rows.forEach((row, index) => { + if (!isValidThreshold(row.threshold)) { + ctx.addIssue({ code: "custom", message: "Enter a whole number from 1 to 100", path: [index, "threshold"] }); + } else if (rows.filter((other) => other.threshold === row.threshold).length > 1) { + ctx.addIssue({ code: "custom", message: "Duplicate threshold", path: [index, "threshold"] }); + } + }); + }) + .optional(), budget_duration: z.string().nullish(), tpm_limit: numericInputSchema, rpm_limit: numericInputSchema, @@ -422,6 +442,7 @@ const TEAM_MEMBER_SETTINGS_FIELDS = [ "team_member_key_duration", "team_member_tpm_limit", "team_member_rpm_limit", + "team_member_max_budget_alert_emails", ] as const; const SEARCH_TOOL_SETTINGS_FIELDS = ["object_permission_search_tools"] as const; @@ -437,6 +458,7 @@ const EMPTY_TEAM_UPDATE_VALUES: TeamUpdateFormValues = { team_member_key_duration: undefined, team_member_tpm_limit: undefined, team_member_rpm_limit: undefined, + team_member_max_budget_alert_emails: [], budget_duration: undefined, tpm_limit: undefined, rpm_limit: undefined, @@ -487,6 +509,7 @@ const toTeamFormValues = (info: TeamInfoRecord, effectiveGuardrails: string[]): team_member_key_duration: info.metadata?.team_member_key_duration, team_member_tpm_limit: info.team_member_budget_table?.tpm_limit, team_member_rpm_limit: info.team_member_budget_table?.rpm_limit, + team_member_max_budget_alert_emails: [...teamMemberBudgetAlertRowsFromMetadata(info.metadata)], budget_duration: info.budget_duration, tpm_limit: info.tpm_limit, rpm_limit: info.rpm_limit, @@ -572,6 +595,11 @@ const TeamInfoView: React.FC = ({ append: appendModelLimit, remove: removeModelLimit, } = useFieldArray({ control: form.control, name: "modelLimits" }); + const { + fields: memberBudgetAlertRows, + append: appendMemberBudgetAlertRow, + remove: removeMemberBudgetAlertRow, + } = useFieldArray({ control: form.control, name: "team_member_max_budget_alert_emails" }); const [teamMemberSettingsOpen, setTeamMemberSettingsOpen] = useState(false); const [searchToolSettingsOpen, setSearchToolSettingsOpen] = useState(false); const [isEditMemberModalVisible, setIsEditMemberModalVisible] = useState(false); @@ -994,6 +1022,15 @@ const TeamInfoView: React.FC = ({ ? { allowed_passthrough_routes: info.metadata.allowed_passthrough_routes } : {}; + const memberBudgetAlertEmails = + values.team_member_max_budget_alert_emails !== undefined + ? teamMemberBudgetAlertEmailsFromRows(values.team_member_max_budget_alert_emails) + : info.metadata?.[TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY]; + const memberBudgetAlertEmailsMetadata = + memberBudgetAlertEmails !== undefined && Object.keys(memberBudgetAlertEmails).length > 0 + ? { [TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY]: memberBudgetAlertEmails } + : {}; + const updateData: any = { team_id: teamId, team_alias: values.team_alias, @@ -1025,6 +1062,7 @@ const TeamInfoView: React.FC = ({ .filter((email: string) => email.length > 0) : values.soft_budget_alerting_emails || [], ...(secretManagerSettings !== undefined ? { secret_manager_settings: secretManagerSettings } : {}), + ...memberBudgetAlertEmailsMetadata, }, ...(values.policies?.length > 0 ? { policies: values.policies } : {}), ...(values.organization_id !== info.organization_id ? { organization_id: values.organization_id ?? null } : {}), @@ -1632,6 +1670,71 @@ const TeamInfoView: React.FC = ({ )} + + + {labelWithHint( + "Budget Alert Thresholds", + "Email each member when their spend reaches a percentage of their team member budget. The member is always notified; add comma-separated addresses to notify others as well. Requires email alerting to be configured on the proxy.", + )} + + {memberBudgetAlertRows.map((row, index) => ( +
+ + {({ ref, value, onChange, ...field }) => ( + ) => + onChange(event.target.value === "" ? null : Number(event.target.value)) + } + placeholder="% of budget" + min={1} + max={100} + step={1} + /> + )} + + + {({ ref, value, ...field }) => ( + + )} + + +
+ ))} + +
@@ -2202,6 +2305,7 @@ const TeamInfoView: React.FC = ({
Key Duration: {info.metadata?.team_member_key_duration || "No Limit"}
TPM Limit: {info.team_member_budget_table?.tpm_limit ?? "No Limit"}
RPM Limit: {info.team_member_budget_table?.rpm_limit ?? "No Limit"}
+
Budget Alert Thresholds: {teamMemberBudgetAlertSummary(info.metadata).join("; ") || "None"}

Router Settings

diff --git a/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.test.ts b/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.test.ts new file mode 100644 index 00000000000..3faa051047b --- /dev/null +++ b/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.test.ts @@ -0,0 +1,94 @@ +import { describe, expect, it } from "vitest"; +import { + isValidThreshold, + teamMemberBudgetAlertEmailsFromRows, + teamMemberBudgetAlertRowsFromMetadata, + teamMemberBudgetAlertSummary, +} from "./teamMemberBudgetAlertEmails"; + +describe("teamMemberBudgetAlertRowsFromMetadata", () => { + it("turns the stored threshold map into rows sorted by threshold", () => { + const metadata = { + team_member_max_budget_alert_emails: { "100": ["finance@example.com", "cto@example.com"], "50": [] }, + }; + expect(teamMemberBudgetAlertRowsFromMetadata(metadata)).toEqual([ + { threshold: 50, emails: "" }, + { threshold: 100, emails: "finance@example.com, cto@example.com" }, + ]); + }); + + it("drops non-numeric thresholds and non-list recipients instead of crashing", () => { + const metadata = { + team_member_max_budget_alert_emails: { fifty: [], "75": "finance@example.com", "90": [1], "100": ["a@b.c"] }, + }; + expect(teamMemberBudgetAlertRowsFromMetadata(metadata)).toEqual([{ threshold: 100, emails: "a@b.c" }]); + }); + + it("drops API-stored thresholds outside 1 to 100 so they never block the form", () => { + const metadata = { + team_member_max_budget_alert_emails: { "0": ["a@b.c"], "50": [], "101": ["a@b.c"] }, + }; + expect(teamMemberBudgetAlertRowsFromMetadata(metadata)).toEqual([{ threshold: 50, emails: "" }]); + }); + + it.each([undefined, null, "50", { team_member_max_budget_alert_emails: "50" }, { soft_budget_alerting_emails: [] }])( + "returns no rows for unrelated or malformed metadata %j", + (metadata) => { + expect(teamMemberBudgetAlertRowsFromMetadata(metadata)).toEqual([]); + }, + ); +}); + +describe("teamMemberBudgetAlertEmailsFromRows", () => { + it("builds the threshold map, splitting, trimming and deduplicating recipients", () => { + expect( + teamMemberBudgetAlertEmailsFromRows([ + { threshold: 50, emails: "" }, + { threshold: 100, emails: " finance@example.com,cto@example.com , finance@example.com, " }, + ]), + ).toEqual({ "50": [], "100": ["finance@example.com", "cto@example.com"] }); + }); + + it("skips rows without a valid threshold", () => { + expect( + teamMemberBudgetAlertEmailsFromRows([ + { threshold: null, emails: "finance@example.com" }, + { threshold: 0, emails: "" }, + { threshold: 101, emails: "" }, + { threshold: 12.5, emails: "" }, + { threshold: 80, emails: "" }, + ]), + ).toEqual({ "80": [] }); + }); + + it("round-trips the stored config", () => { + const stored = { team_member_max_budget_alert_emails: { "50": [], "100": ["finance@example.com"] } }; + expect(teamMemberBudgetAlertEmailsFromRows(teamMemberBudgetAlertRowsFromMetadata(stored))).toEqual( + stored.team_member_max_budget_alert_emails, + ); + }); +}); + +describe("isValidThreshold", () => { + it.each([ + [1, true], + [50, true], + [100, true], + [0, false], + [101, false], + [33.3, false], + [null, false], + ])("treats %s as valid=%s", (threshold, valid) => { + expect(isValidThreshold(threshold)).toBe(valid); + }); +}); + +describe("teamMemberBudgetAlertSummary", () => { + it("states that the member is always notified and lists extra recipients", () => { + expect( + teamMemberBudgetAlertSummary({ + team_member_max_budget_alert_emails: { "100": ["finance@example.com"], "50": [] }, + }), + ).toEqual(["50%: member", "100%: member, finance@example.com"]); + }); +}); diff --git a/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.ts b/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.ts new file mode 100644 index 00000000000..36d5ddbac02 --- /dev/null +++ b/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.ts @@ -0,0 +1,57 @@ +export const TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY = "team_member_max_budget_alert_emails" as const; + +export interface TeamMemberBudgetAlertRow { + readonly threshold: number | null; + readonly emails: string; +} + +export type TeamMemberBudgetAlertEmails = Readonly>; + +const isEmailList = (value: unknown): value is readonly string[] => + Array.isArray(value) && value.every((email) => typeof email === "string"); + +const splitEmails = (emails: string): readonly string[] => + Array.from( + new Set( + emails + .split(",") + .map((email) => email.trim()) + .filter((email) => email.length > 0), + ), + ); + +const THRESHOLD_MIN = 1; +const THRESHOLD_MAX = 100; + +export const isValidThreshold = (threshold: number | null): threshold is number => { + const isWholeNumber = threshold !== null && Number.isInteger(threshold); + return isWholeNumber && threshold >= THRESHOLD_MIN && threshold <= THRESHOLD_MAX; +}; + +export const teamMemberBudgetAlertRowsFromMetadata = (metadata: unknown): readonly TeamMemberBudgetAlertRow[] => { + if (typeof metadata !== "object" || metadata === null) return []; + const config: unknown = (metadata as Record)[TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY]; + if (typeof config !== "object" || config === null || Array.isArray(config)) return []; + return Object.entries(config as Record) + .flatMap(([key, emails]) => { + const threshold = Number(key); + return /^\d+$/.test(key) && isValidThreshold(threshold) && isEmailList(emails) + ? [{ threshold, emails: emails.join(", ") }] + : []; + }) + .sort((a, b) => (a.threshold ?? 0) - (b.threshold ?? 0)); +}; + +export const teamMemberBudgetAlertEmailsFromRows = ( + rows: readonly TeamMemberBudgetAlertRow[], +): TeamMemberBudgetAlertEmails => + Object.fromEntries( + rows + .filter((row) => isValidThreshold(row.threshold)) + .map((row) => [String(row.threshold), splitEmails(row.emails)]), + ); + +export const teamMemberBudgetAlertSummary = (metadata: unknown): readonly string[] => + teamMemberBudgetAlertRowsFromMetadata(metadata).map((row) => + row.emails.length > 0 ? `${row.threshold}%: member, ${row.emails}` : `${row.threshold}%: member`, + ); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 1e0aed46923..6d46af04730 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -32957,6 +32957,8 @@ export interface components { cache_creation_input_token_cost_above_1hr?: number | null; /** Cache Creation Input Token Cost Above 200K Tokens */ cache_creation_input_token_cost_above_200k_tokens?: number | null; + /** Cache Creation Input Token Cost Above 200K Tokens Batches */ + cache_creation_input_token_cost_above_200k_tokens_batches?: number | null; /** Cache Creation Input Token Cost Above 272K Tokens */ cache_creation_input_token_cost_above_272k_tokens?: number | null; /** Cache Creation Input Token Cost Above 272K Tokens Batches */ @@ -46750,6 +46752,8 @@ export interface components { cache_creation_input_token_cost_above_1hr?: number | null; /** Cache Creation Input Token Cost Above 200K Tokens */ cache_creation_input_token_cost_above_200k_tokens?: number | null; + /** Cache Creation Input Token Cost Above 200K Tokens Batches */ + cache_creation_input_token_cost_above_200k_tokens_batches?: number | null; /** Cache Creation Input Token Cost Above 272K Tokens */ cache_creation_input_token_cost_above_272k_tokens?: number | null; /** Cache Creation Input Token Cost Above 272K Tokens Batches */