Merge remote-tracking branch 'origin/main' into litellm_mcp_persistent_upstream_session

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

# Conflicts:
#	tests/integration/mcp/test_mcp_lifecycle.py
This commit is contained in:
Devin AI 2026-09-23 22:57:33 +00:00
commit 8007215599
344 changed files with 30333 additions and 4975 deletions

View file

@ -3222,7 +3222,7 @@ workflows:
cron: "17 0,6,12,18 * * *"
filters:
branches:
only: litellm_internal_staging
only: main
jobs: *migration_jobs
integration:
unless: << pipeline.parameters.run_migration_tests >>

View file

@ -11,7 +11,7 @@ run_full() {
[ -n "${CIRCLE_PULL_REQUEST:-}" ] || run_full "not a pull request"
candidate_bases="main litellm_internal_staging litellm_oss_staging"
candidate_bases="${PATH_FILTER_BASE_BRANCH:-main}"
merge_base=""
for base in $candidate_bases; do
git fetch --quiet origin "$base" 2>/dev/null || continue

292
.circleci/tests.yml Normal file
View file

@ -0,0 +1,292 @@
version: 2.1
commands:
wait_for_service:
parameters:
url:
type: string
timeout:
type: string
default: "60"
steps:
- run:
name: "Wait for << parameters.url >>"
command: |
TIMEOUT=<< parameters.timeout >>
URL="<< parameters.url >>"
ELAPSED=0
echo "Waiting up to ${TIMEOUT}s for ${URL} ..."
if echo "$URL" | grep -q '^tcp://'; then
HOST=$(echo "$URL" | sed 's|tcp://||' | cut -d: -f1)
PORT=$(echo "$URL" | sed 's|tcp://||' | cut -d: -f2)
while ! bash -c "echo > /dev/tcp/$HOST/$PORT" 2>/dev/null; do
sleep 2; ELAPSED=$((ELAPSED+2))
if [ "$ELAPSED" -ge "$TIMEOUT" ]; then echo "Timed out"; exit 1; fi
done
else
while ! curl -sf --max-time 5 "$URL" > /dev/null 2>&1; do
sleep 2; ELAPSED=$((ELAPSED+2))
if [ "$ELAPSED" -ge "$TIMEOUT" ]; then echo "Timed out"; exit 1; fi
done
fi
echo "Service ready after ${ELAPSED}s"
install_uv:
steps:
- run:
name: Install uv (pinned 0.10.9)
command: |
curl -LsSf -o /tmp/uv-install.sh https://astral.sh/uv/0.10.9/install.sh
echo "7fc46e39cb97290b57169c0c813a17970585ac519139f19006453c99b5f2f45f /tmp/uv-install.sh" | sha256sum -c -
env UV_NO_MODIFY_PATH=1 sh /tmp/uv-install.sh
rm -f /tmp/uv-install.sh
echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$BASH_ENV"
export PATH="$HOME/.local/bin:$PATH"
install_rust:
steps:
- run:
name: Install Rust (rustup 1.28.2, toolchain 1.98.0)
command: |
case "$(uname -m)" in
x86_64)
RUSTUP_TRIPLE=x86_64-unknown-linux-gnu
RUSTUP_SHA256=20a06e644b0d9bd2fbdbfd52d42540bdde820ea7df86e92e533c073da0cdd43c
;;
aarch64)
RUSTUP_TRIPLE=aarch64-unknown-linux-gnu
RUSTUP_SHA256=e3853c5a252fca15252d07cb23a1bdd9377a8c6f3efa01531109281ae47f841c
;;
*)
echo "install_rust: unsupported architecture $(uname -m)" >&2
exit 1
;;
esac
curl -sSLf -o /tmp/rustup-init \
"https://static.rust-lang.org/rustup/archive/1.28.2/${RUSTUP_TRIPLE}/rustup-init"
echo "${RUSTUP_SHA256} /tmp/rustup-init" | sha256sum -c -
chmod +x /tmp/rustup-init
/tmp/rustup-init -y --no-modify-path --profile minimal --default-toolchain 1.98.0
rm -f /tmp/rustup-init
echo 'export PATH="$HOME/.cargo/bin:$PATH"' >> "$BASH_ENV"
export PATH="$HOME/.cargo/bin:$PATH"
rustc --version
cargo --version
install_codecov_cli:
steps:
- run:
name: Install Codecov CLI (pinned v11.3.1)
command: |
curl -sSLf -o /tmp/codecov https://cli.codecov.io/v11.3.1/linux/codecov
curl -sSLf -o /tmp/codecov.SHA256SUM https://cli.codecov.io/v11.3.1/linux/codecov.SHA256SUM
[ "$(cat /tmp/codecov.SHA256SUM)" = "ca1d64196d2d34771084afe76ea657d581bf628e31d993ff8e52ea09cc88a56d codecov" ]
(cd /tmp && sha256sum -c codecov.SHA256SUM)
chmod +x /tmp/codecov
mkdir -p "$HOME/.local/bin"
mv /tmp/codecov "$HOME/.local/bin/codecov"
setup_litellm_enterprise_pip:
steps:
- run:
name: "Install local version of litellm-enterprise"
command: |
uv run --no-sync python -c "import litellm_enterprise; print('litellm-enterprise OK:', litellm_enterprise.__file__)"
setup_test_deps:
steps:
- checkout
- install_uv
- install_rust
- restore_cache:
keys:
- v3-integration-uv-cache-{{ checksum "uv.lock" }}
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- setup_litellm_enterprise_pip
- save_cache:
paths:
- ~/.cache/uv
key: v3-integration-uv-cache-{{ checksum "uv.lock" }}
- run:
name: Generate Prisma client
command: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
skip_unless_relevant:
parameters:
category:
type: string
default: backend
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
steps:
- run:
name: "Skip job when no << parameters.category >>-relevant files changed"
command: |
export CIRCLE_PULL_REQUEST="${CIRCLE_PULL_REQUEST:-<< parameters.pull_request_url >>}"
export PATH_FILTER_BASE_BRANCH="<< parameters.base_ref >>"
[ -n "$PATH_FILTER_BASE_BRANCH" ] || unset PATH_FILTER_BASE_BRANCH
bash .circleci/scripts/path_filter.sh << parameters.category >>
start_postgres:
parameters:
db_name:
type: string
default: circle_test
image:
type: string
default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26
steps:
- run:
name: Start PostgreSQL
command: |
docker run -d \
--name postgres-db \
-e POSTGRES_USER=postgres \
-e POSTGRES_PASSWORD=postgres \
-e POSTGRES_DB=<< parameters.db_name >> \
-p 5432:5432 \
<< parameters.image >>
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
start_redis:
steps:
- run:
name: Start Redis
command: |
docker run -d \
--name redis-cache \
-p 6379:6379 \
redis:7-alpine@sha256:7aec734b2bb298a1d769fd8729f13b8514a41bf90fcdd1f38ec52267fbaa8ee6
- wait_for_service:
url: tcp://localhost:6379
timeout: "60"
jobs:
unit:
parameters:
tests_path:
type: string
default: tests/unit
flag:
type: string
default: unit
shards:
type: integer
default: 6
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
parallelism: << parameters.shards >>
environment:
LITELLM_LOCAL_MODEL_COST_MAP: "True"
steps:
- setup_test_deps
- skip_unless_relevant:
base_ref: << parameters.base_ref >>
pull_request_url: << parameters.pull_request_url >>
- run:
name: "Run << parameters.tests_path >> shard"
no_output_timeout: 20m
command: |
mkdir -p test-results/<< parameters.flag >>
mapfile -t files < <(find << parameters.tests_path >> -name 'test_*.py' | sort | circleci tests split --split-by=timings --timings-type=filename)
if [ "${#files[@]}" -eq 0 ]; then echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.tests_path >> files; nothing to run"; exit 0; fi
set +e
uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml
status=$?
set -e
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from the shard; passing"; exit 0; fi
exit "$status"
- install_codecov_cli
- run:
name: Upload coverage
when: always
command: |
[ -f coverage.xml ] || { echo "no coverage.xml produced; skipping upload"; exit 0; }
codecov upload-process --disable-search -f coverage.xml -F << parameters.flag >> -C "$CIRCLE_SHA1" -n "<< parameters.flag >>-${CIRCLE_NODE_INDEX}-${CIRCLE_BUILD_NUM}" --git-service github
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
- store_artifacts:
path: coverage.xml
documentation:
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- setup_test_deps
- run:
name: Checkout litellm-docs
command: rm -rf docs/my-website && git clone --depth 1 https://github.com/BerriAI/litellm-docs.git docs/my-website
- run:
name: Run documentation validation
command: |
uv run --no-sync python ./tests/documentation_tests/test_env_keys.py
uv run --no-sync python ./tests/documentation_tests/test_router_settings.py
uv run --no-sync python ./tests/documentation_tests/test_api_docs.py
uv run --no-sync python ./tests/documentation_tests/test_circular_imports.py
integration:
parameters:
suite:
type: string
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- setup_test_deps
- skip_unless_relevant:
base_ref: << parameters.base_ref >>
pull_request_url: << parameters.pull_request_url >>
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
- start_redis
- run:
name: Run owned integration contracts
command: bash .circleci/scripts/run_integration.sh << parameters.suite >>
no_output_timeout: 15m
- run:
name: Stop owned database and Redis
when: always
command: |
mkdir -p test-results/integration-<< parameters.suite >>
docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true
docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
workflows:
tests:
when: (pipeline.event.name == "push" and pipeline.git.branch == "main") or pipeline.event.name == "api" or (pipeline.event.name == "pull_request" and (pipeline.event.github.pull_request.base.ref == "main" or pipeline.event.github.pull_request.base.ref starts-with "litellm_"))
jobs:
- unit:
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- documentation
- integration:
name: integration-<< matrix.suite >>
matrix:
parameters:
suite: [sdk]
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>

View file

@ -8,7 +8,6 @@
#
# Protected branches (always allowed):
# - main
# - litellm_internal_staging
# - dependabot/*
# - gh-readonly-queue/*
#
@ -22,7 +21,7 @@ ZERO_OID_SHA256="000000000000000000000000000000000000000000000000000000000000000
ALLOWED_TYPES="feature|bugfix|hotfix|release|chore"
BRANCH_PATTERN="^(${ALLOWED_TYPES})/.+"
PROTECTED_NAMES="main litellm_internal_staging"
PROTECTED_NAMES="main"
PROTECTED_PREFIXES="dependabot/ gh-readonly-queue/"
is_protected() {
@ -78,8 +77,7 @@ if [ -n "$invalid" ]; then
chore/bump-deps
hotfix/auth-bypass
Protected (always allowed): main, litellm_internal_staging,
dependabot/*, gh-readonly-queue/*.
Protected (always allowed): main, dependabot/*, gh-readonly-queue/*.
See https://conventional-branch.github.io/

View file

@ -10,12 +10,13 @@ test_paths:
paths:
- tests/rust-python-harness
- reason: >-
What is left of the caching suite in tests/local_testing that runs nowhere. Every job that
globs that directory either deselects it (local_testing_part1 and part2 carry `-k "... and
not caching and not cache"`) or keeps only another keyword (langfuse, router, assistants),
and no job names these files the way redis_caching_unit_tests names test_dual_cache.py.
The gap was eight files and 118 tests when measured 2026-08-20; the five keyless ones now
run in the caching-local shard, leaving these three. Measured 2026-08-21 with no provider
Live-provider caching cases in tests/local_testing that remain outside CI. Jobs that
glob that directory either deselect them (local_testing_part1 and part2 carry `-k "... and
not caching and not cache"`) or keep only another keyword (langfuse, router, assistants).
Separately, test-redis-compat.yml selects two IAM cluster authentication tests in
test_caching.py by node ID. It does not run that file's other tests.
The gap was eight files and 118 tests when measured 2026-08-20; the five keyless files now
run in the caching-local shard, leaving live cases in these three. Measured 2026-08-21 with no provider
credentials and no Redis: test_caching.py needs both (37 of 65 fail without them),
test_disk_cache_unit_tests.py needs OPENAI_API_KEY for 2 of its 4, and
test_gcs_cache_unit_tests.py needs GCS credentials for all 4. They want the keyless/live

15
.github/merge-smoke-tests.json vendored Normal file
View file

@ -0,0 +1,15 @@
{
"cases": {
"CHAT-JSON": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
"CHAT-TEXT-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
"CHAT-TOOL-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
"MODEL-ALLOW": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_allows_listed_model_for_key",
"MODEL-DENY": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]",
"COST-EXPLICIT": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
"COST-ZERO": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
"LOG-CONTENT-ON": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
"LOG-CONTENT-OFF": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
"CALLBACK-SUCCESS": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
"CALLBACK-FAILURE": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
}
}

View file

@ -1,162 +1,27 @@
<!-- The whole description's target audience is humans, not AI agents: write it in plain, simple,
everyday engineering language, extremely parsable and readable at a glance. This goes double for
the TLDR, User Flow, and Caveats sections -->
<!-- Plain English please. Describe the change the way you would explain it to a teammate who has not seen the code: what it does and why, not which functions, files, or tables it touches -->
## TLDR
## What's the problem?
<!-- Fill in the bullets below and keep each one short and concrete: one line per bullet, roughly 10 words max -->
## What's the solution?
Problem this solves:
<!-- The approach, in a sentence or two -->
- <blah>
- ...
## How does it fix it?
How it solves it:
<!-- What actually changes so the problem can't happen anymore -->
- <blah>
- ...
## How does the product experience change?
## User Flow
<!-- What a user could do or see before, and what they can do or see after. If nothing user-facing changes, say so -->
<!-- Two ordered lists, Before and After, walking the same end user through the same task, written strictly from that user's seat
Read the linked issue, ticket, or customer thread first so the flow reflects the real application and the routes its users actually hit; don't invent a generic scenario
Lead each list with one plain sentence saying where the flow fails (Before) or succeeds (After), then number the steps
Every step is something the user does or observes: the HTTP method and full URL they hit, what they sent, and what visibly came back (status code, error text, the shape of an ID). UI steps name the page URL and what is on screen
No LiteLLM internals: never name functions, files, DB tables, config classes, hooks, callbacks, or code paths. "The upload hands back an ID that looks like OpenAI's own `file-abc123` instead of the scrambled one the gateway returned" is right, "no managed-file row was registered" is wrong
Keep the two lists step-for-step identical until they diverge, so the changed step is obvious
If the bug had a security or authorization consequence, end each list with what another user could or could no longer do
Regenerate this section whenever new commits change the PR's behavior, so it never describes an older revision
## What caveats are there, if any?
Example:
Before: a developer whose app streams chat completions gets no token counts back, so their cost dashboard reads zero
1. They send POST https://litellm-domain/v1/chat/completions with `"stream": true` and no `stream_options`
2. The last SSE chunk arrives with `"usage": null`, so their app records 0 prompt and 0 completion tokens
3. They open https://litellm-domain/ui/?page=logs and see the request logged at $0 spend
After: the same request comes back with real token counts, so the dashboard shows real spend
1. The proxy admin sets `always_include_stream_usage: true` and restarts the proxy
2. The developer sends the same POST https://litellm-domain/v1/chat/completions with `"stream": true` and no `stream_options`
3. The last SSE chunk now carries a `usage` object with real prompt and completion token counts
4. https://litellm-domain/ui/?page=logs shows that request at non-zero spend
-->
## Relevant issues
<!-- e.g., "Fixes #000" -->
## Affected release
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1", and add the `backport-stable` label so the fix is cherry-picked onto the rc line before the stable is tagged. Leave the section blank otherwise -->
<!-- Major ones only: behavior that breaks on purpose, migrations that lock tables, auth changes, known gaps you did not fix. Write "None" if there are none -->
## Linear ticket
<!-- if you are an internal contributor, add "Resolves " followed by the Linear ticket e.g., "Resolves LIT-1234" to link the Linear ticket to the GitHub PR. If you don't have one, leave the section blank rather than guessing -->
<!-- Internal contributors: Resolves LIT-1234. Otherwise leave blank -->
## Pre-Submission checklist
**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/<your_test_file>.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)
## Delays in PR merge?
If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slack (#pr-review)](https://join.slack.com/t/litellmossslack/shared_invite/zt-3o7nkuyfr-p_kbNJj8taRfXGgQI1~YyA).
## Screenshots / Proof of Fix
<!-- Include screenshots, screen recordings, or command (e.g., curl) + output demonstrating that your changes work as expected
The proof must be completely e2e with no mocks, using actual LLM calls costing real $$$ if applicable. `pytest` commands are not enough
Show ONLY the latest run: capture Before at the merge base and After at the PR's current tip, and when new commits change behavior, replace this whole section with the fresh run instead of stacking it on top of older ones. The run must be up to date. As soon as a new commit is made and it makes this PR description's after sha stale (it's no longer tip of PR), you must re-run the QA
Structure the section exactly as below: Before and After one heading level below this section, each naming the commit hash it was captured at, one lower-level heading per case inside each, the same case names in the same order on both sides, and numbered steps (command, observed output) under every case, never loose prose; shared setup (config, payloads) goes above Before, and with a single case, drop the case headings and number the steps directly
### Before (<hash>)
#### <case 1>
1. ...
2. ...
#### <case 2>
1. ...
### After (<hash>)
#### <case 1>
1. ...
2. ...
#### <case 2>
1. ...
For bug fixes: Before shows the reproduction, After shows the same steps passing
For new features: Before shows the capability missing, After shows it working end-to-end
If the change applies to all three LLM endpoints (/v1/responses, /v1/chat/completions, /v1/messages), make each endpoint its own case, not just one
For UI changes: before/after screenshots under the same headings
If the main use case runs through a coding tool like Claude Code or Codex, drive that tool interactively the way the user does (never `claude -p`, `codex exec`, or curl on its own) and embed before/after screenshots of its pane under the same headings; curl replays and headless runs can follow as extra cases, never as the only proof -->
## Type
<!-- Select the type of Pull Request -->
<!-- Keep only the necessary ones -->
🆕 New Feature
🐛 Bug Fix
🧹 Refactoring
📖 Documentation
🚄 Infrastructure
✅ Test
## Caveats (if any)
<!-- Group caveats under severity subheadings (### Severe, ### High, ### Medium, ### Low), with
short bullet points inside each, just like the TLDR: one line per bullet, roughly 10 words max
Call out known limitations, follow-up work, or anything a reviewer should watch out for
Include only the tiers that have caveats; drop the empty ones
- Severe: inherent to what the PR deliberately ships, there even when the code works as intended:
it can degrade or take down a running deployment (e.g. a slow or table-locking boot migration),
rewrite data by design, break an existing workflow on purpose, or change auth behavior. An
operator must plan around it before rollout
- High: an unintended hole: a correctness, security, data-loss, or backward-compatibility bug,
unsafe to ship as is
- Medium: a real gap someone can hit, but with a workaround or a narrow blast radius
- Low: anything else worth noting: naming, cleanup, an edge case nobody hits
Nest bullets as deep as helps: hierarchy beats one long line when it makes things clearer to a
human reader
If you assumed something instead of testing it, e.g. "only reproduces with X on" or "no
user-observable behavior difference", list it here too with what breaks if it is wrong
Leave this section empty if there are none -->
## QA runbook
<!-- Only needed when your PR edits tests/e2e; delete this section otherwise
For each e2e test you added or changed, list the manual steps a reviewer can follow to reproduce it by hand against a live proxy, mapping 1:1 to what the test asserts: one top-level bullet per test giving its pytest node id followed by what it proves in plain words, then a nested "- [ ]" checklist where each item is a concrete action (route, request body, expected response) and the final item is the sanity-check step shown in the examples. Note environment prerequisites (provider credentials, config flags) and any nuances a manual run will hit. See PRs #32914 and #32963 for full examples
Example checklists:
- tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py::TestKeyRateLimits::test_rpm_limit_blocks_over_limit - a key allowed 2 requests a minute serves exactly 2 and refuses the 3rd
- [ ] Generate a limited key: curl -X POST http://localhost:4000/key/generate -H "Authorization: Bearer sk-1234" -d '{"rpm_limit": 2}'
- [ ] Send three /v1/chat/completions requests with that key inside one minute
- [ ] Expect the first two to return 200 and the third to return 429 naming the rpm limit
- [ ] Sanity check: this test makes sense to add and is not hand-wavey (e.g., assert actual expected spend instead of just spend > 0) or potentially flaky
- tests/e2e/management/test_management_e2e.py::TestModelRoutes::test_model_create_appears_in_ui - a deployment created through the API shows up on the Admin UI models page
- [ ] POST /model/new with the master key, a bedrock model, and aws_region_name (needs STORE_MODEL_IN_DB=True and AWS credentials)
- [ ] Open http://localhost:4000/ui/?page=models and expect a deployment row showing the returned model id
- [ ] Sanity check: this test makes sense to add and is not hand-wavey (e.g., assert actual expected spend instead of just spend > 0) or potentially flaky
-->
## Final Attestation
- [ ] The tests check the right things, including the edge cases, and regressions in the respective real-world customer use-cases are not possible after this PR
## How did you test this?
<!-- What you ran against a live proxy and what came back. Commands with output or before/after screenshots work well. Unit tests alone are not enough -->

493
.github/scripts/run_merge_smoke.py vendored Normal file
View file

@ -0,0 +1,493 @@
#!/usr/bin/env python3
"""Merge smoke harness: bounded checks run inside a loopback-only Linux network namespace."""
# ruff: noqa: T201 # CLI harness: stdout/stderr lines are the reported result
from __future__ import annotations
import argparse
import contextlib
import http.client
import json
import os
import secrets
import signal
import socket
import subprocess
import sys
import time
from collections import Counter
from collections.abc import Sequence
from dataclasses import dataclass, field
from pathlib import Path
from types import MappingProxyType
from typing import Final, NoReturn, TextIO, cast
import pytest
EXPECTED_CASES: Final = (
"CHAT-JSON",
"CHAT-TEXT-STREAM",
"CHAT-TOOL-STREAM",
"MODEL-ALLOW",
"MODEL-DENY",
"COST-EXPLICIT",
"COST-ZERO",
"LOG-CONTENT-ON",
"LOG-CONTENT-OFF",
"CALLBACK-SUCCESS",
"CALLBACK-FAILURE",
)
@dataclass(frozen=True, slots=True)
class CheckResult:
ok: bool
detail: str = ""
@dataclass(slots=True)
class _Args:
command: str = ""
no_child: bool = False
expect: str = ""
litellm_bin: str | None = None
lite_bin: str | None = None
diagnostics_dir: str = ""
ready_deadline: float = 120.0
shutdown_deadline: float = 20.0
poll_interval: float = 0.5
manifest: str = ""
rootdir: str | None = None
def fail(reason: str) -> NoReturn:
print(f"merge-smoke: FAIL {reason}", file=sys.stderr)
sys.exit(1)
def ok(step: str) -> None:
print(f"merge-smoke: OK {step}")
def tail(path: Path, lines: int = 20) -> str:
try:
return "\n".join(path.read_text(errors="replace").splitlines()[-lines:])
except OSError as exc:
return f"<cannot read {path}: {exc}>"
def cmd_verify_isolation(args: _Args) -> int:
if os.geteuid() == 0:
fail("verify-isolation must run unprivileged (geteuid()==0)")
try:
socket.create_connection(("192.0.2.1", 9), timeout=3)
except OSError as exc:
print(f"external connect blocked as expected: errno={exc.errno} {exc}")
else:
fail("external TCP connect to 192.0.2.1:9 succeeded; namespace is not isolated")
listener: Final = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
listener.bind(("127.0.0.1", 0))
listener.listen(1)
port: Final = cast(int, listener.getsockname()[1])
client: Final = socket.create_connection(("127.0.0.1", port), timeout=5)
accepted: Final = listener.accept()
accepted[0].close()
client.close()
listener.close()
print(f"loopback connect ok on 127.0.0.1:{port}")
if not args.no_child:
proc: Final = subprocess.run(
[sys.executable, str(Path(__file__).resolve()), "verify-isolation", "--no-child"],
timeout=30,
capture_output=True,
text=True,
)
if proc.returncode != 0:
fail(f"child process did not inherit isolation: {proc.stderr.strip()}")
print("child process inherits isolation")
ok("verify-isolation")
return 0
def cmd_interpreter(args: _Args) -> int:
print(sys.version)
print(sys.executable)
actual: Final = f"{sys.version_info.major}.{sys.version_info.minor}"
if actual != args.expect:
fail(f"interpreter is {actual}, expected {args.expect}")
ok(f"interpreter {actual}")
return 0
def _run_cli(argv: Sequence[str], label: str) -> CheckResult:
try:
proc: Final = subprocess.run(list(argv), timeout=120, capture_output=True, text=True)
except subprocess.TimeoutExpired:
return CheckResult(ok=False, detail=f"{label} timed out after 120s")
sys.stdout.write(proc.stdout)
sys.stderr.write(proc.stderr)
if proc.returncode != 0:
return CheckResult(ok=False, detail=f"{label} exited {proc.returncode}")
return CheckResult(ok=True)
def cmd_cli(args: _Args) -> int:
venv_bin: Final = Path(sys.executable).parent
litellm_bin: Final = Path(args.litellm_bin) if args.litellm_bin else venv_bin / "litellm"
lite_bin: Final = Path(args.lite_bin) if args.lite_bin else venv_bin / "lite"
commands: Final = (
("import litellm", [sys.executable, "-c", "import litellm"]),
("litellm --version", [str(litellm_bin), "--version"]),
("lite version", [str(lite_bin), "version"]),
)
for label, argv in commands:
result = _run_cli(argv, label)
if not result.ok:
fail(result.detail)
ok(label)
return 0
def _free_port() -> int:
sock: Final = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
sock.bind(("127.0.0.1", 0))
port: Final = cast(int, sock.getsockname()[1])
sock.close()
return port
_CONFIG_TEMPLATE: Final = """model_list:
- model_name: smoke-model
litellm_params:
model: openai/smoke-model
api_base: http://127.0.0.1:9/v1
api_key: synthetic-key
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
"""
def _listen_inode(port: int) -> str | None:
target: Final = f"{port:04X}"
for table in ("/proc/net/tcp", "/proc/net/tcp6"):
try:
rows = Path(table).read_text().splitlines()[1:]
except OSError:
continue
for row in rows:
cols = row.split()
if len(cols) > 9 and cols[3] == "0A" and cols[1].rsplit(":", 1)[-1] == target:
return cols[9]
return None
def _ancestors(pid: int) -> frozenset[int]:
chain: Final[set[int]] = set()
pending: Final[list[int]] = [pid]
while pending:
current = pending.pop()
if current <= 0 or current in chain:
continue
chain.add(current)
try:
stat = Path(f"/proc/{current}/stat").read_text()
except OSError:
continue
pending.append(int(stat.rpartition(")")[2].split()[1]))
return frozenset(chain)
def _socket_owner_pid(inode: str) -> int | None:
for proc_dir in Path("/proc").iterdir():
if not proc_dir.name.isdigit():
continue
fd_dir = proc_dir / "fd"
try:
for fd in fd_dir.iterdir():
try:
if os.readlink(fd) == f"socket:[{inode}]":
return int(proc_dir.name)
except OSError:
continue
except OSError:
continue
return None
def _verify_port_owner(port: int, proc: subprocess.Popen[bytes]) -> CheckResult:
inode: Final = _listen_inode(port)
if inode is None:
return CheckResult(ok=False, detail=f"no LISTEN socket found for port {port} in /proc/net/tcp")
owner: Final = _socket_owner_pid(inode)
if owner is None:
return CheckResult(ok=False, detail=f"no process owns the listen socket inode {inode} for port {port}")
if owner != proc.pid and proc.pid not in _ancestors(owner):
return CheckResult(
ok=False, detail=f"port {port} owned by pid {owner} outside the launched process group {proc.pid}"
)
if proc.poll() is not None:
return CheckResult(ok=False, detail=f"proxy exited with code {proc.returncode} after readiness")
return CheckResult(ok=True)
def cmd_proxy_startup(args: _Args) -> int:
diagnostics: Final = Path(args.diagnostics_dir)
diagnostics.mkdir(parents=True, exist_ok=True)
venv_bin: Final = Path(sys.executable).parent
litellm_bin: Final = Path(args.litellm_bin) if args.litellm_bin else venv_bin / "litellm"
port: Final = _free_port()
master_key: Final = "sk-smoke-" + secrets.token_hex(16)
config_path: Final = diagnostics / "config.yaml"
config_path.write_text(_CONFIG_TEMPLATE)
log_path: Final = diagnostics / "proxy.log"
result_path: Final = diagnostics / "result.json"
outcome: Final[dict[str, object]] = {
"port": port,
"time_to_ready_s": None,
"shutdown_s": None,
"readiness": None,
"outcome": "failed",
}
log_file: Final = log_path.open("w")
env: Final = {
**os.environ,
"LITELLM_MASTER_KEY": master_key,
"LITELLM_LOCAL_MODEL_COST_MAP": "True",
}
started: Final = time.monotonic()
proc: Final = subprocess.Popen(
[str(litellm_bin), "--config", str(config_path), "--host", "127.0.0.1", "--port", str(port)],
stdout=log_file,
stderr=subprocess.STDOUT,
start_new_session=True,
env=env,
)
body: str | None = None
last_status: int | None = None
while time.monotonic() - started < args.ready_deadline:
if proc.poll() is not None:
log_file.close()
result_path.write_text(json.dumps(outcome))
fail(f"proxy exited early with code {proc.returncode}\n{tail(log_path)}")
try:
conn = http.client.HTTPConnection("127.0.0.1", port, timeout=5)
conn.request("GET", "/health/readiness")
resp = conn.getresponse()
last_status = resp.status
candidate = resp.read().decode()
conn.close()
except (http.client.HTTPException, ConnectionError, OSError):
time.sleep(args.poll_interval)
continue
if last_status == 200:
body = candidate
break
time.sleep(args.poll_interval)
outcome["time_to_ready_s"] = round(time.monotonic() - started, 3)
if body is None:
_terminate(proc, log_file)
result_path.write_text(json.dumps(outcome))
detail = f"last status {last_status}" if last_status is not None else "no response"
fail(f"readiness not reached within {args.ready_deadline}s ({detail})\n{tail(log_path)}")
outcome["readiness"] = body
try:
readiness = cast(object, json.loads(body))
except json.JSONDecodeError:
readiness = None
if readiness != {"status": "healthy", "db": "Not connected"}:
_terminate(proc, log_file)
result_path.write_text(json.dumps(outcome))
fail(f"unexpected readiness body: {body}")
owner_check: Final = _verify_port_owner(port, proc)
if not owner_check.ok:
_terminate(proc, log_file)
result_path.write_text(json.dumps(outcome))
fail(owner_check.detail)
shutdown_started: Final = time.monotonic()
os.killpg(proc.pid, signal.SIGTERM)
try:
proc.wait(timeout=args.shutdown_deadline)
except subprocess.TimeoutExpired:
os.killpg(proc.pid, signal.SIGKILL)
proc.wait(timeout=10)
outcome["shutdown_s"] = round(time.monotonic() - shutdown_started, 3)
log_file.close()
result_path.write_text(json.dumps(outcome))
fail(f"forced kill after {args.shutdown_deadline}s\n{tail(log_path)}")
outcome["shutdown_s"] = round(time.monotonic() - shutdown_started, 3)
try:
os.killpg(proc.pid, 0)
except ProcessLookupError:
pass
else:
os.killpg(proc.pid, signal.SIGKILL)
log_file.close()
result_path.write_text(json.dumps(outcome))
fail("process group survived SIGTERM")
log_file.close()
outcome["outcome"] = "ok"
result_path.write_text(json.dumps(outcome))
ok(f"proxy-startup ready={outcome['time_to_ready_s']}s shutdown={outcome['shutdown_s']}s")
return 0
def _terminate(proc: subprocess.Popen[bytes], log_file: TextIO) -> None:
with contextlib.suppress(ProcessLookupError):
os.killpg(proc.pid, signal.SIGTERM)
try:
proc.wait(timeout=10)
except subprocess.TimeoutExpired:
with contextlib.suppress(ProcessLookupError):
os.killpg(proc.pid, signal.SIGKILL)
with contextlib.suppress(subprocess.TimeoutExpired):
proc.wait(timeout=10)
log_file.close()
def _load_manifest(path: Path) -> MappingProxyType[str, str]:
def no_duplicates(pairs: list[tuple[object, object]]) -> dict[object, object]:
seen: dict[object, object] = {}
for key, value in pairs:
if key in seen:
raise ValueError(f"duplicate key in manifest: {key}")
seen[key] = value
return seen
raw_value: object = cast(object, json.loads(path.read_text(), object_pairs_hook=no_duplicates))
if not isinstance(raw_value, dict):
raise ValueError("manifest must be an object")
loaded: Final = cast(dict[object, object], raw_value)
cases_value: object = loaded.get("cases")
if not isinstance(cases_value, dict):
raise ValueError("manifest must be an object with a 'cases' object")
cases_any: Final = cast(dict[object, object], cases_value)
cases: Final = {k: v for k, v in cases_any.items() if isinstance(k, str) and isinstance(v, str)}
if len(cases) != len(cases_any):
raise ValueError("manifest 'cases' must map string ids to string node ids")
return MappingProxyType(cases)
@dataclass(slots=True, eq=False)
class _Recorder:
collect_failed: list[str] = field(default_factory=list)
collected: tuple[str, ...] = ()
reports: dict[str, list[tuple[str, str, bool]]] = field(default_factory=dict)
def pytest_collectreport(self, report: pytest.CollectReport) -> None:
if report.failed:
self.collect_failed.append(report.nodeid)
def pytest_collection_finish(self, session: pytest.Session) -> None:
self.collected = tuple(item.nodeid for item in session.items)
def pytest_runtest_logreport(self, report: pytest.TestReport) -> None:
self.reports.setdefault(report.nodeid, []).append((report.when, report.outcome, hasattr(report, "wasxfail")))
def cmd_pytest(args: _Args) -> int:
try:
cases: Final = _load_manifest(Path(args.manifest))
except (OSError, ValueError, json.JSONDecodeError) as exc:
fail(f"manifest invalid: {exc}")
if tuple(cases) != EXPECTED_CASES:
fail(f"manifest case ids must be exactly {list(EXPECTED_CASES)} in order, got {list(cases)}")
node_ids: Final = tuple(cases.values())
if len(set(node_ids)) != len(node_ids):
fail("manifest node ids are not unique")
argv: Final = [
*node_ids,
"-p",
"no:cacheprovider",
"-p",
"no:xdist",
"-p",
"no:rerunfailures",
"-p",
"no:randomly",
"-rA",
"-q",
*(["--rootdir", args.rootdir] if args.rootdir else []),
]
recorder: Final = _Recorder()
code: Final = pytest.main(argv, plugins=[recorder])
name_of: Final = MappingProxyType({node_id: case_id for case_id, node_id in cases.items()})
problems: Final[list[str]] = []
if code != 0:
problems.append(f"pytest exit code {code}")
for failed_id in recorder.collect_failed:
problems.append(f"collection failed: {name_of.get(failed_id, failed_id)}")
expected: Final = Counter(node_ids)
collected: Final = Counter(recorder.collected)
for node_id in expected - collected:
problems.append(f"missing case {name_of[node_id]} ({node_id})")
for node_id in collected - expected:
problems.append(f"unexpected test collected: {node_id}")
for node_id, count in collected.items():
if count > 1:
problems.append(f"duplicated test id: {node_id}")
if len(recorder.collected) != len(EXPECTED_CASES):
problems.append(f"collected {len(recorder.collected)} tests, expected {len(EXPECTED_CASES)}")
rows: Final[list[tuple[str, bool]]] = []
for case_id, node_id in cases.items():
reports = recorder.reports.get(node_id, [])
case_ok = (
bool(reports)
and all(outcome == "passed" and not wasxfail for _, outcome, wasxfail in reports)
and {when for when, _, _ in reports} >= {"setup", "call", "teardown"}
)
rows.append((case_id, case_ok))
if not reports:
problems.append(f"{case_id} ({node_id}) produced no runtest reports")
continue
for when, outcome, wasxfail in reports:
if outcome != "passed":
problems.append(f"{case_id} ({node_id}) {when} outcome={outcome}")
if wasxfail:
problems.append(f"{case_id} ({node_id}) {when} was xfail/xpass")
missing_phases = {"setup", "call", "teardown"} - {when for when, _, _ in reports}
for phase in sorted(missing_phases):
problems.append(f"{case_id} ({node_id}) missing {phase} report")
for case_id, passed in rows:
print(f"{case_id} {'PASS' if passed else 'FAIL'} {cases[case_id]}")
if problems:
for problem in problems:
print(f"merge-smoke: {problem}", file=sys.stderr)
fail("pytest verdict failed")
ok("pytest 11 cases")
return 0
def main() -> int:
parser: Final = argparse.ArgumentParser(description=__doc__)
subs: Final = parser.add_subparsers(dest="command", required=True)
p_iso: Final = subs.add_parser("verify-isolation")
p_iso.add_argument("--no-child", action="store_true")
p_interp: Final = subs.add_parser("interpreter")
p_interp.add_argument("--expect", required=True)
p_cli: Final = subs.add_parser("cli")
p_cli.add_argument("--litellm-bin", default=None)
p_cli.add_argument("--lite-bin", default=None)
p_proxy: Final = subs.add_parser("proxy-startup")
p_proxy.add_argument("--diagnostics-dir", required=True)
p_proxy.add_argument("--litellm-bin", default=None)
p_proxy.add_argument("--ready-deadline", type=float, default=120)
p_proxy.add_argument("--shutdown-deadline", type=float, default=20)
p_proxy.add_argument("--poll-interval", type=float, default=0.5)
p_test: Final = subs.add_parser("pytest")
p_test.add_argument("--manifest", required=True)
p_test.add_argument("--rootdir", default=None)
args: Final = parser.parse_args(namespace=_Args())
handlers: Final = {
"verify-isolation": cmd_verify_isolation,
"interpreter": cmd_interpreter,
"cli": cmd_cli,
"proxy-startup": cmd_proxy_startup,
"pytest": cmd_pytest,
}
return handlers[args.command](args)
if __name__ == "__main__":
sys.exit(main())

View file

@ -214,7 +214,7 @@ def main(
native_module: Final = load_native_module(native_path)
native_module_loads: Final = native_module is not None
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
native_size_limit: Final = 35_000_000
native_size_limit: Final = 40_000_000
native_size_within_limit: Final = native_member.file_size <= native_size_limit
validations: Final = (
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),

View file

@ -4,8 +4,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
permissions:

View file

@ -4,13 +4,10 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read

View file

@ -67,12 +67,18 @@ jobs:
# further up the stack are modified. The suppression is scoped to this one
# file/rule pair via SARIF post-filtering so every other callsite of
# py/weak-sensitive-data-hashing in the repository continues to be analyzed.
- name: Filter SARIF (OCI sha256)
# The same query fires on the HIBP k-anonymity lookup in
# litellm/proxy/auth/password_policy.py, where the password's SHA-1 is only
# a lookup key into the haveibeenpwned range API (the protocol mandates
# SHA-1) and the digest itself never leaves the proxy beyond its first 5
# characters.
- name: Filter SARIF (OCI sha256, HIBP sha1)
if: matrix.language == 'python'
uses: advanced-security/filter-sarif@2da736ff05ef065cb2894ac6892e47b5eac2c3c0 # v1.1
with:
patterns: |
-litellm/llms/oci/common_utils.py:py/weak-sensitive-data-hashing
-litellm/proxy/auth/password_policy.py:py/weak-sensitive-data-hashing
input: sarif-results/python.sarif
output: sarif-results/python.sarif

View file

@ -4,7 +4,6 @@ on:
push:
branches:
- main
- litellm_internal_staging
paths:
- "litellm/**"
- "tests/benchmarks/**"
@ -17,7 +16,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
paths:
- "litellm/**"
- "tests/benchmarks/**"

View file

@ -4,8 +4,6 @@ on: # zizmor: ignore[dangerous-triggers] runs the base branch's code only; the P
pull_request_target:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
permissions:

View file

@ -1,49 +0,0 @@
name: Create Daily Staging Branch
on:
schedule:
- cron: "0 0,12 * * *" # Runs every 12 hours at midnight and noon UTC
workflow_dispatch: # Allow manual trigger
jobs:
create-staging-branch:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- name: Create daily staging branch
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
BRANCH_NAME="litellm_oss_staging_$(date +'%m_%d_%Y')"
echo "Creating branch: $BRANCH_NAME"
if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then
echo "Branch $BRANCH_NAME already exists. Skipping creation."
exit 0
fi
MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha')
gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent
echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA"
create-internal-dev-branch:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- name: Create internal dev branch
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
BRANCH_NAME="litellm_internal_dev_$(date +'%m_%d_%Y')"
echo "Creating branch: $BRANCH_NAME"
if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then
echo "Branch $BRANCH_NAME already exists. Skipping creation."
exit 0
fi
MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha')
gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent
echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA"

View file

@ -4,8 +4,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "uv.lock"

View file

@ -4,7 +4,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_branch
- "litellm_**"
paths:

View file

@ -4,8 +4,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
schedule:
- cron: "23 6 * * *"

View file

@ -1,6 +1,6 @@
name: Publish basedpyright base counts
# Every commit on main or litellm_internal_staging can become a future merge-base.
# Every commit on main can become a future merge-base.
# Publishing its per-rule basedpyright counts as an artifact lets
# scripts/type_check_gate.py download them in seconds instead of paying a
# 60-110s second basedpyright pass on every fresh worktree or moved merge-base.
@ -11,7 +11,6 @@ on:
push:
branches:
- main
- litellm_internal_staging
workflow_dispatch:
inputs:
ref:

View file

@ -4,13 +4,10 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
@ -83,6 +80,9 @@ jobs:
- name: test_e2e_changed_gate
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_e2e_changed_gate.py tests/code_coverage_tests/test_e2e_idp_stack.py
- name: Check merge smoke harness
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_merge_smoke.py
- name: router_code_coverage
run: uv run --no-sync python ./tests/code_coverage_tests/router_code_coverage.py

View file

@ -4,8 +4,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
permissions:

View file

@ -7,8 +7,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
concurrency:

View file

@ -6,8 +6,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
concurrency:

View file

@ -7,13 +7,10 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}

View file

@ -4,8 +4,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
permissions:

95
.github/workflows/test-merge-smoke.yml vendored Normal file
View file

@ -0,0 +1,95 @@
name: Merge smoke checks
on:
pull_request:
branches: [main, litellm_internal_staging, litellm_oss_staging, "litellm_**"]
workflow_dispatch:
permissions:
contents: read
concurrency:
group: merge-smoke-${{ github.event.pull_request.number || github.run_id }}
cancel-in-progress: true
jobs:
dashboard-build:
name: Dashboard build
runs-on: ubuntu-24.04
timeout-minutes: 30
steps:
- name: Checkout
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Build the dashboard stage
run: docker build --target ui-builder -f Dockerfile .
core-checks:
name: Core checks (Python ${{ matrix.python-version }})
runs-on: ubuntu-24.04
timeout-minutes: 30
strategy:
fail-fast: false
matrix:
python-version: ["3.10", "3.11", "3.12", "3.13", "3.14"]
env:
LITELLM_LOCAL_MODEL_COST_MAP: "True"
steps:
- name: Checkout
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: ${{ matrix.python-version }}
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install dependencies
run: .github/scripts/uv_sync_with_retries.sh --frozen --extra proxy --extra cli --group dev --group proxy-dev --python ${{ matrix.python-version }}
- name: Create the loopback-only network namespace
run: |
sudo ip netns add smoke
sudo ip netns exec smoke ip link set lo up
cat > "${RUNNER_TEMP}/in-netns" <<'WRAP'
#!/usr/bin/env bash
set -euo pipefail
exec sudo --preserve-env=LITELLM_LOCAL_MODEL_COST_MAP ip netns exec smoke setpriv --reuid "$(id -u)" --regid "$(id -g)" --init-groups -- env HOME="${HOME}" PATH="${PATH}" "$@"
WRAP
chmod +x "${RUNNER_TEMP}/in-netns"
echo "IN_NETNS=${RUNNER_TEMP}/in-netns" >> "${GITHUB_ENV}"
- name: Verify namespace isolation
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py verify-isolation
- name: Verify interpreter version
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py interpreter --expect ${{ matrix.python-version }}
- name: Import and CLI checks
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py cli
- name: Proxy startup check
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py proxy-startup --diagnostics-dir "${RUNNER_TEMP}/smoke-diagnostics"
- name: Run curated smoke cases
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py pytest --manifest .github/merge-smoke-tests.json
- name: Upload smoke diagnostics
if: always()
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: merge-smoke-diagnostics-py${{ matrix.python-version }}
path: ${{ runner.temp }}/smoke-diagnostics
if-no-files-found: ignore
- name: Remove the network namespace
if: always()
run: sudo ip netns delete smoke

View file

@ -4,13 +4,10 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
workflow_dispatch:
permissions:

View file

@ -4,13 +4,12 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "litellm/_redis.py"
- "litellm/_redis_credential_provider.py"
- "tests/test_litellm/test_redis.py"
- "tests/local_testing/test_caching.py"
- "tests/test_litellm/caching/test_redis_connection_pool.py"
- ".github/workflows/test-redis-compat.yml"
- "pyproject.toml"
@ -28,6 +27,9 @@ jobs:
name: "redis-py ${{ matrix.redis-version }}"
runs-on: ubuntu-latest
timeout-minutes: 15
permissions:
contents: read
id-token: write
strategy:
fail-fast: false
@ -57,7 +59,7 @@ jobs:
- name: Install dependencies
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra extra_proxy --extra semantic-router
- name: Pin redis-py to the matrix version
env:
@ -66,12 +68,33 @@ jobs:
uv pip install "redis==${REDIS_VERSION:?}"
uv run --no-sync python -c "import redis; assert redis.__version__ == '${REDIS_VERSION:?}', redis.__version__; print('redis-py', redis.__version__)"
- name: Build Redis for cluster authentication tests
run: |
curl --fail --location --retry 3 https://download.redis.io/releases/redis-7.2.16.tar.gz -o "$RUNNER_TEMP/redis-7.2.16.tar.gz"
echo "960a8ec15e34ff40e57ff16837b26b33bd81f2da6d24497bb63de532a323a18e $RUNNER_TEMP/redis-7.2.16.tar.gz" | sha256sum --check
tar -xzf "$RUNNER_TEMP/redis-7.2.16.tar.gz" -C "$RUNNER_TEMP"
make -C "$RUNNER_TEMP/redis-7.2.16" -j2 MALLOC=libc OPTIMIZATION=-O1 redis-server
echo "$RUNNER_TEMP/redis-7.2.16/src" >> "$GITHUB_PATH"
- name: Run redis unit tests
run: |
redis-server --version
uv run --no-sync pytest \
tests/test_litellm/test_redis.py \
tests/test_litellm/caching/test_redis_connection_pool.py \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \
--tb=short -vv \
--reruns 2 \
--reruns-delay 1 \
--durations=20
--durations=20 \
--cov=./litellm --cov-report=xml:coverage-redis.xml
- name: Upload Redis coverage
if: matrix.redis-version == '5.3.1'
uses: codecov/codecov-action@75cd11691c0faa626561e295848008c8a7dddffe # v5.5.4
with:
use_oidc: true
files: coverage-redis.xml
flags: redis-compat
fail_ci_if_error: false

View file

@ -29,8 +29,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "litellm-rust/**"

View file

@ -4,8 +4,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
permissions:

View file

@ -9,8 +9,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "terraform/litellm/aws/**"

View file

@ -8,8 +8,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "terraform/provider/**"

View file

@ -4,13 +4,10 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read

View file

@ -4,13 +4,10 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read

View file

@ -4,13 +4,10 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
workflow_dispatch:
permissions:
@ -114,6 +111,7 @@ jobs:
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/messages
tests/test_litellm/embeddings
tests/test_litellm/ocr
tests/test_litellm/passthrough
tests/test_litellm/rag

View file

@ -6,8 +6,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "vscode-extension/**"

View file

@ -2,12 +2,10 @@ name: GitHub Actions Security Analysis
on:
push:
branches: [main, litellm_internal_staging]
branches: [main]
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
concurrency:

View file

@ -39,7 +39,7 @@ Same applies for filing bug reports and feature requests, with .github/ISSUE_TEM
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just leave the section blank
Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, or the main use case runs through a headful agentic coding tool like Claude Code or Codex, drive that surface yourself and embed your own before and after screenshots of it in the PR (the Admin UI page, or what the coding tool shows), next to an ordered list of the URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, and what fields to fill out so a reviewer can reproduce it
Never use `pytest` commands or the like as the answer to "How did you test this?". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, or the main use case runs through a headful agentic coding tool like Claude Code or Codex, drive that surface yourself and embed your own before and after screenshots of it in the PR (the Admin UI page, or what the coding tool shows), next to an ordered list of the URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, and what fields to fill out so a reviewer can reproduce it
If you ever write any human-facing text (pull requests, issues, commit messages, discussion posts, github comments, release notes, docs, etc.), always follow these guidelines to sound less AI-y:
- don't use emojis

193
litellm-rust/Cargo.lock generated
View file

@ -88,6 +88,45 @@ version = "1.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "03918c3dbd7701a85c6b9887732e2921175f26c350b4563841d0958c21d57e6d"
[[package]]
name = "asn1-rs"
version = "0.7.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b7f43a50ac4fdca5df8e885c21b835997f0a1cdee65494a6847694a98652d9d8"
dependencies = [
"asn1-rs-derive",
"asn1-rs-impl",
"displaydoc",
"nom",
"num-traits",
"rusticata-macros",
"thiserror 2.0.19",
"time",
]
[[package]]
name = "asn1-rs-derive"
version = "0.6.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3109e49b1e4909e9db6515a30c633684d68cdeaa252f215214cb4fa1a5bfee2c"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.119",
"synstructure 0.13.2",
]
[[package]]
name = "asn1-rs-impl"
version = "0.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7b18050c2cd6fe86c3a76584ef5e0baf286d038cda203eb6223df2cc413565f7"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.119",
]
[[package]]
name = "assert-json-diff"
version = "2.0.2"
@ -810,7 +849,7 @@ version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3"
dependencies = [
"bit-vec",
"bit-vec 0.8.0",
]
[[package]]
@ -819,6 +858,15 @@ version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7"
[[package]]
name = "bit-vec"
version = "0.9.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b71798fca2c1fe1086445a7258a4bc81e6e49dcd24c8d0dd9a1e57395b603f51"
dependencies = [
"serde",
]
[[package]]
name = "bitflags"
version = "1.3.2"
@ -1120,6 +1168,15 @@ version = "0.4.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "338089f42c427b86394a5ee60ff321da23a5c89c9d89514c829687b26359fcff"
[[package]]
name = "crc32c"
version = "0.6.8"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3a47af21622d091a8f0fb295b88bc886ac74efcc613efc19f5d0b21de5c89e47"
dependencies = [
"rustc_version",
]
[[package]]
name = "crc32fast"
version = "1.5.1"
@ -1378,6 +1435,20 @@ dependencies = [
"thiserror 2.0.19",
]
[[package]]
name = "der-parser"
version = "10.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "07da5016415d5a3c4dd39b11ed26f915f52fc4e0dc197d87908bc916e51bc1a6"
dependencies = [
"asn1-rs",
"displaydoc",
"nom",
"num-bigint 0.4.8",
"num-traits",
"rusticata-macros",
]
[[package]]
name = "deranged"
version = "0.5.8"
@ -2589,21 +2660,6 @@ dependencies = [
"wasm-bindgen",
]
[[package]]
name = "jsonwebtoken"
version = "11.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e75fe14a82d81e5f5af639997db37d8b96045938a7ac6ab18cdbe1c7467e05e1"
dependencies = [
"base64 0.22.1",
"getrandom 0.2.17",
"js-sys",
"serde",
"serde_json",
"signature",
"zeroize",
]
[[package]]
name = "lazy_static"
version = "1.5.0"
@ -3119,6 +3175,7 @@ dependencies = [
"tokio",
"tokio-tungstenite",
"url",
"veil",
"wiremock",
]
@ -3143,10 +3200,11 @@ version = "0.1.0"
dependencies = [
"aws-sdk-kms",
"base64 0.22.1",
"futures-util",
"google-cloud-auth",
"google-cloud-kms-v1",
"jsonwebtoken",
"litellm-core-utils",
"litellm-python-compat",
"litellm-secrets-aws",
"litellm-secrets-azure",
"litellm-secrets-cyberark",
@ -3179,6 +3237,7 @@ dependencies = [
"litellm-tracing",
"rstest",
"serde_json",
"tempfile",
"thiserror 2.0.19",
"tokio",
"veil",
@ -3215,10 +3274,12 @@ dependencies = [
"litellm-tracing",
"moka",
"percent-encoding",
"rcgen",
"reqwest 0.12.28",
"rstest",
"serde",
"serde_json",
"tempfile",
"thiserror 2.0.19",
"tokio",
"veil",
@ -3230,6 +3291,7 @@ name = "litellm-secrets-google"
version = "0.1.0"
dependencies = [
"base64 0.22.1",
"crc32c",
"google-cloud-auth",
"google-cloud-gax",
"google-cloud-kms-v1",
@ -3255,7 +3317,6 @@ version = "0.1.0"
dependencies = [
"litellm-core-utils",
"litellm-secrets-types",
"moka",
"rstest",
"rustify",
"rustify_derive",
@ -3274,6 +3335,7 @@ name = "litellm-secrets-types"
version = "0.1.0"
dependencies = [
"litellm-auth-types",
"moka",
"rstest",
"serde",
"serde_json",
@ -3588,6 +3650,15 @@ dependencies = [
"libc",
]
[[package]]
name = "oid-registry"
version = "0.8.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "12f40cff3dde1b6087cc5d5f5d4d65712f34016a03ed60e9c08dcc392736b5b7"
dependencies = [
"asn1-rs",
]
[[package]]
name = "once_cell"
version = "1.21.4"
@ -3721,6 +3792,16 @@ version = "0.2.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2ee67f1008b1ba2321834326597b8e186293b049a023cdef258527550b9935b4"
[[package]]
name = "pem"
version = "4.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d354a98a3d1251555de99e8fdd8afda05573c31b82f59063a7b0a29b5527f120"
dependencies = [
"base64 0.23.1",
"serde_core",
]
[[package]]
name = "percent-encoding"
version = "2.3.2"
@ -3899,7 +3980,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4b45fcc2344c680f5025fe57779faef368840d0bd1f42f216291f0dc4ace4744"
dependencies = [
"bit-set",
"bit-vec",
"bit-vec 0.8.0",
"bitflags 2.13.1",
"num-traits",
"rand 0.9.5",
@ -4288,6 +4369,20 @@ dependencies = [
"crossbeam-utils",
]
[[package]]
name = "rcgen"
version = "0.14.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8774e05a7d0de114588e6a28fe7e71694b82614ed569d86d8b389dfbc98b8ad8"
dependencies = [
"pem",
"ring",
"rustls-pki-types",
"time",
"x509-parser",
"yasna",
]
[[package]]
name = "redis"
version = "1.7.0"
@ -4570,6 +4665,15 @@ dependencies = [
"semver",
]
[[package]]
name = "rusticata-macros"
version = "4.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "faf0c4a6ece9950b9abdb62b1cfcf2a68b3b67a10ba445b3bb85be2a293d0632"
dependencies = [
"nom",
]
[[package]]
name = "rustify"
version = "0.7.0"
@ -5046,15 +5150,6 @@ dependencies = [
"libc",
]
[[package]]
name = "signature"
version = "2.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de"
dependencies = [
"rand_core 0.6.4",
]
[[package]]
name = "simd-adler32"
version = "0.3.10"
@ -6356,6 +6451,24 @@ version = "0.6.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4"
[[package]]
name = "x509-parser"
version = "0.18.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d43b0f71ce057da06bc0851b23ee24f3f86190b07203dd8f567d0b706a185202"
dependencies = [
"asn1-rs",
"data-encoding",
"der-parser",
"lazy_static",
"nom",
"oid-registry",
"ring",
"rusticata-macros",
"thiserror 2.0.19",
"time",
]
[[package]]
name = "xmlparser"
version = "0.13.6"
@ -6368,6 +6481,16 @@ version = "0.8.18"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "aee1b19627c7c60102ab80d3a9cbe18de90bfe03bfa6c3715447681f0e8c8af6"
[[package]]
name = "yasna"
version = "0.6.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b5f6765e852b9b4dc8e2a76843e4d64d1cea8e79bcde0b6901aea8e7c7f08282"
dependencies = [
"bit-vec 0.9.1",
"time",
]
[[package]]
name = "yoke"
version = "0.8.3"
@ -6437,20 +6560,6 @@ name = "zeroize"
version = "1.9.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e"
dependencies = [
"zeroize_derive",
]
[[package]]
name = "zeroize_derive"
version = "1.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3c50655cbb0fe3fc43170059e702f1ce5e19b84cec58dc87b037a09935c2f328"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.119",
]
[[package]]
name = "zerotrie"

View file

@ -7,6 +7,7 @@ repository.workspace = true
autotests = false
[dependencies]
litellm-secrets.workspace = true
litellm-types.workspace = true
litellm-core-utils.workspace = true
litellm-host.workspace = true
@ -36,7 +37,6 @@ url.workspace = true
veil.workspace = true
[dev-dependencies]
litellm-secrets.workspace = true
litellm-auth-gcp.workspace = true
litellm-llms = { workspace = true, features = ["test-support"] }
rstest.workspace = true

View file

@ -1,11 +1,9 @@
use litellm_auth::{InputSource, SecretValue, Sourced};
use litellm_llms::base_llm::{
inference::secrets::Secrets,
ocr::{
handler::OcrClient,
transformation::{OcrConnection, OcrCredentialInputs, PreparedOcrRequest},
},
use litellm_llms::base_llm::ocr::{
handler::OcrClient,
transformation::{OcrConnection, OcrCredentialInputs, PreparedOcrRequest},
};
use litellm_secrets::source::Secrets;
use super::provider_config::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, ResolvedOcrRequest};

View file

@ -11,7 +11,6 @@ use litellm_http::{
HttpClientPool, HttpSettings, Resolution,
media::{PublicDnsResolver, UrlPolicy},
};
use litellm_llms::base_llm::inference::secrets::{SecretSource, Secrets};
use litellm_llms::base_llm::ocr::{
error::Error as OcrError,
handler::OcrClient,
@ -20,6 +19,7 @@ use litellm_llms::base_llm::ocr::{
BaseOcrConfig, LiteLLMOcrResponse, OCR_RESPONSE_MAX_BYTES, OcrTransportConfig,
},
};
use litellm_secrets::source::SecretSource;
use rstest::rstest;
use serde_json::{Value, json};
@ -32,27 +32,27 @@ use super::{
use crate::ocr::route::{LocalOcrHost, OcrOp, OcrOpResult, ocr_machine};
struct RecordingSecretSource {
names: Arc<Mutex<Vec<&'static str>>>,
names: Arc<Mutex<Vec<String>>>,
values: &'static [(&'static str, &'static str)],
api_base: String,
}
impl SecretSource for RecordingSecretSource {
fn resolve<'a>(
fn get_secret_str<'a>(
&'a self,
names: &'a [&'static str],
) -> BoxFuture<'a, Result<Secrets, litellm_secrets::Error>> {
*self.names.lock().unwrap() = names.to_vec();
let values = self.values;
let api_base = self.api_base.clone();
name: &'a str,
) -> BoxFuture<'a, Result<Option<litellm_secrets::SecretValue>, litellm_secrets::Error>> {
self.names.lock().unwrap().push(name.to_owned());
Box::pin(async move {
Ok(Arc::new(move |name: &str| match name {
"MISTRAL_AZURE_API_BASE" => Some(api_base.clone()),
_ => values
Ok(match name {
"MISTRAL_AZURE_API_BASE" => Some(self.api_base.clone()),
_ => self
.values
.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| value.to_string()),
}) as Secrets)
}
.map(litellm_secrets::SecretValue::new))
})
}
}
@ -289,7 +289,7 @@ async fn ocr_client_uses_the_injected_http_pool_configuration() {
UrlPolicy::default(),
VertexAuth::default(),
OcrSettings::default(),
Arc::new(litellm_llms::base_llm::inference::secrets::EnvironmentSecrets),
Arc::new(litellm_secrets::source::EnvironmentSecrets::default()),
)
.unwrap();
crate::ocr::client::perform(&client, wire_request("mistral/model", &base, json!({})))

View file

@ -27,7 +27,9 @@ pub use execution::{
pub use fork_gate::RuntimeAlreadyStarted;
pub use gil::{release_count, release_gil};
pub use handle::{Execution, ExecutionBody, ExecutionStep};
pub use marshal::{Pythonized, from_py, from_py_argument, panic_to_pyerr, to_py};
pub use marshal::{
Pythonized, from_py, from_py_argument, json_loads, json_object_field, panic_to_pyerr, to_py,
};
/// Starts the interpreter and imports the standard modules the tests share, once, so
/// parallel test threads never race a first import of `asyncio`.

View file

@ -4,6 +4,7 @@ use std::panic::{AssertUnwindSafe, catch_unwind};
use pyo3::exceptions::PyValueError;
use pyo3::panic::PanicException;
use pyo3::prelude::*;
use pyo3::types::PyBytes;
use serde::Serialize;
use serde::de::DeserializeOwned;
@ -32,6 +33,19 @@ where
.map_err(PyErr::from)
}
pub fn json_object_field(py: Python<'_>, document: &str, name: &str) -> PyResult<Py<PyAny>> {
py.import("json")?
.call_method1("loads", (document,))?
.call_method1("get", (name,))
.map(Bound::unbind)
}
pub fn json_loads(py: Python<'_>, document: &[u8]) -> PyResult<Py<PyAny>> {
py.import("json")?
.call_method1("loads", (PyBytes::new(py, document),))
.map(Bound::unbind)
}
pub struct Pythonized<T>(pub T);
impl<'py, T> IntoPyObject<'py> for Pythonized<T>

View file

@ -1 +0,0 @@
pub mod secrets;

View file

@ -1,19 +0,0 @@
use std::sync::Arc;
use futures_util::future::BoxFuture;
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
use litellm_secrets::Error;
pub type Secrets = Arc<dyn Lookup + Send + Sync>;
pub trait SecretSource: Send + Sync {
fn resolve<'a>(&'a self, names: &'a [&'static str]) -> BoxFuture<'a, Result<Secrets, Error>>;
}
pub struct EnvironmentSecrets;
impl SecretSource for EnvironmentSecrets {
fn resolve<'a>(&'a self, _names: &'a [&'static str]) -> BoxFuture<'a, Result<Secrets, Error>> {
Box::pin(async { Ok(Arc::new(ProcessEnvironment) as Secrets) })
}
}

View file

@ -2,6 +2,5 @@ pub mod anthropic_messages;
pub mod audio_transcription;
pub mod base_model_iterator;
pub mod chat;
pub mod inference;
pub mod ocr;
pub mod responses;

View file

@ -13,7 +13,6 @@ use litellm_http::{
use serde::{Serialize, de::DeserializeOwned};
use serde_json::Value;
use crate::base_llm::inference::secrets::SecretSource;
use crate::base_llm::ocr::{
error::Error,
settings::OcrSettings,
@ -22,6 +21,7 @@ use crate::base_llm::ocr::{
PreparedOcrRequest, decode_request_value, decode_response,
},
};
use litellm_secrets::source::SecretSource;
/// The route's view of one call, handed to provider code that has to reach the
/// caller's hooks mid-flight (guardrails on the outgoing body, raw response events).
@ -95,7 +95,7 @@ impl OcrClient {
document_fetcher: MediaFetcher::for_test(document_http),
vertex_auth: VertexAuth::default(),
settings: OcrSettings::default(),
secrets: Arc::new(crate::base_llm::inference::secrets::EnvironmentSecrets),
secrets: Arc::new(litellm_secrets::source::EnvironmentSecrets::default()),
}
}

View file

@ -7,6 +7,7 @@ use litellm_core_utils::{
settings::ProcessEnvironment,
};
use litellm_http::outbound::{OutboundRequest, RequestSigner};
use litellm_secrets::source::Secrets;
use serde::{
Deserialize, Serialize,
de::{DeserializeOwned, IntoDeserializer},
@ -14,13 +15,10 @@ use serde::{
use serde_json::{Map, Value};
use serde_with::serde_as;
use crate::base_llm::{
inference::secrets::Secrets,
ocr::{
error::Error,
handler::{CallHooks, OcrClient, read_response_bytes, transform_request_body},
settings::OcrSettings,
},
use crate::base_llm::ocr::{
error::Error,
handler::{CallHooks, OcrClient, read_response_bytes, transform_request_body},
settings::OcrSettings,
};
pub const OCR_RESPONSE_MAX_BYTES: usize = 64 * 1024 * 1024;

View file

@ -45,7 +45,7 @@ litellm-core-utils.workspace = true
litellm-auth-gcp.workspace = true
litellm-http.workspace = true
litellm-llms.workspace = true
litellm-secrets = { workspace = true, features = ["aws"] }
litellm-secrets = { workspace = true, features = ["aws", "azure", "google", "hashicorp", "cyberark"] }
litellm-secrets-types.workspace = true
litellm-types.workspace = true
litellm-host-python.workspace = true
@ -55,6 +55,7 @@ pyo3-async-runtimes.workspace = true
reqwest.workspace = true
redis = { version = "1.7.0", features = ["tls-rustls"] }
serde_json.workspace = true
veil.workspace = true
thiserror.workspace = true
tokio = { workspace = true, features = ["rt", "sync"] }
url.workspace = true

View file

@ -1,5 +1,10 @@
Native OCR uses `SecretSource` with `EnvironmentSecrets`, preserving process-environment reads. Readable Python secret managers still make OCR decline to the existing Python implementation. `ResolvedSecrets` and the separate `secret_manager_binding()` snapshot are inactive foundations for a later rollout
Native OCR uses `litellm_secrets::source::SecretSource`. Built-in secret managers resolve to retained Rust backends. Custom Python managers and overrides keep the callback path. Readable managers still require the Rust secret-manager binding to be enabled
Cache and secret-manager catalog entries remain Python-only, including when `LITELLM_RUST=1`. The new cache runtime is not connected to SDK or gateway caching
The shared proxy initializer captures native configuration without loading the extension or doing native I/O. `_SecretManagerRuntime.from_client` constructs a backend on first use and keeps its handle on the Python client. The secret-manager dispatcher selects Python or Rust through `catalog.py`. Native reads call that handle; Rust routes extract the backend directly. Configuration changes replace the handle, while calls already bound to the previous backend keep using it. Handles cannot be reused after fork. Directly constructed LiteLLM managers are adapted on first native use. Manually supplied SDK clients keep their Python behavior because their credentials cannot be inferred safely. Provider implementations contain no bridge registration
Retention describes ownership and lifetime. `callbacks-legacy-python::PublicCall` owns Python references for one call to preserve identity. A native cache or secret-manager handle owns shared Rust state across calls to preserve connection pools and caches. Both use existing `Py<T>` and shared Rust ownership, with execution and GIL transitions handled by `litellm-host-python`
Cache and secret-manager catalog entries remain Python-only, including when `LITELLM_RUST=1`. This wiring does not change rollout policy
OCR provider requests use the shared `litellm-http` pool. AWS and Google secret-manager SDK clients keep their SDK transports, which do not yet inherit the pool's proxy, TLS, certificate, timeout, or observability configuration. Preserve those SDK transports and configure them equivalently instead of forcing them through reqwest

View file

@ -29,6 +29,7 @@ pub(super) enum CacheBinding {
#[pyclass(frozen, name = "_ResponseCacheRuntime")]
pub(crate) struct ResolvedCache {
binding: CacheBinding,
guard: Option<super::facade::FacadeGuard>,
pid: u32,
}
@ -36,10 +37,24 @@ impl ResolvedCache {
pub(super) fn new(binding: CacheBinding) -> Self {
Self {
binding,
guard: None,
pid: std::process::id(),
}
}
pub(super) fn with_guard(mut self, guard: super::facade::FacadeGuard) -> Self {
self.guard = Some(guard);
self
}
pub(super) fn native_service(&self) -> PyResult<Option<NativeResponseCache>> {
self.check_process()?;
Ok(match &self.binding {
CacheBinding::Native(service) => Some(service.clone()),
_ => None,
})
}
fn check_process(&self) -> PyResult<()> {
if matches!(self.binding, CacheBinding::Native(_)) && self.pid != std::process::id() {
return Err(PyRuntimeError::new_err(
@ -70,6 +85,43 @@ impl ResolvedCache {
#[pymethods]
impl ResolvedCache {
#[staticmethod]
pub(crate) fn from_selected(cache: &Bound<'_, PyAny>) -> PyResult<Self> {
let py = cache.py();
let binding = if cache.is_none() {
CacheBinding::Disabled
} else if let Ok(handle) = cache.extract::<PyRef<'_, super::handle::CacheTestHandle>>() {
CacheBinding::Native(handle.service()?)
} else if let Some(service) = super::facade::resolve(py, cache)? {
CacheBinding::Native(service)
} else if let Some(runtime) = cache
.getattr_opt("_native_cache")?
.filter(|value| !value.is_none())
{
let resolved = runtime
.getattr("native")?
.extract::<PyRef<'_, ResolvedCache>>()?;
match resolved.native_service()? {
Some(service) => {
if !resolved
.guard
.as_ref()
.is_some_and(|guard| guard.matches(py, cache).unwrap_or(false))
{
return Err(RustBridgeDeclined::new_err(
"native cache runtime no longer matches its facade",
));
}
CacheBinding::Native(service)
}
None => CacheBinding::PythonCallback(PythonCallback::new(cache.clone().unbind())),
}
} else {
CacheBinding::PythonCallback(PythonCallback::new(cache.clone().unbind()))
};
Ok(Self::new(binding))
}
#[staticmethod]
fn from_cache(cache: &Bound<'_, PyAny>) -> PyResult<Self> {
let config = match NativeCacheConfig::project(cache)? {
@ -80,7 +132,13 @@ impl ResolvedCache {
};
let backend = cache.getattr("cache")?;
let service = activate(cache.py(), &backend, config)?;
Ok(Self::new(CacheBinding::Native(service)))
let resolved = Self::new(CacheBinding::Native(service.clone()));
Ok(
match super::facade::FacadeGuard::capture(cache.py(), cache, &service) {
Ok(guard) => resolved.with_guard(guard),
Err(_) => resolved,
},
)
}
#[getter]
@ -323,6 +381,9 @@ impl ResolvedCache {
if let CacheBinding::PythonCallback(callback) = &self.binding {
callback.traverse(&visit)?;
}
if let Some(guard) = &self.guard {
guard.traverse(visit)?;
}
Ok(())
}
}

View file

@ -472,7 +472,7 @@ impl FacadeGuard {
})
}
fn matches(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult<bool> {
pub(super) fn matches(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult<bool> {
if !self.outer.matches(py, facade)? {
return Ok(false);
}

View file

@ -20,9 +20,7 @@ use pyo3::{
types::PyDict,
};
pub(crate) use self::{
binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheTestResolver,
};
pub(crate) use self::{binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheResolver};
fn cache_error(error: Error) -> PyErr {
match error {

View file

@ -1,19 +1,14 @@
use pyo3::{PyTraverseError, PyVisit, prelude::*};
use super::{
binding::{CacheBinding, ResolvedCache},
callback::PythonCallback,
facade,
handle::CacheTestHandle,
};
use super::binding::ResolvedCache;
#[pyclass(frozen, name = "_CacheTestResolver")]
pub(crate) struct CacheTestResolver {
#[pyclass(frozen, name = "_CacheResolver")]
pub(crate) struct CacheResolver {
namespace: Py<PyAny>,
}
#[pymethods]
impl CacheTestResolver {
impl CacheResolver {
#[new]
fn new(namespace: Py<PyAny>) -> Self {
Self { namespace }
@ -21,16 +16,7 @@ impl CacheTestResolver {
pub(crate) fn resolve(&self, py: Python<'_>) -> PyResult<ResolvedCache> {
let object = self.namespace.bind(py).getattr("cache")?;
let binding = if object.is_none() {
CacheBinding::Disabled
} else if let Ok(handle) = object.extract::<PyRef<'_, CacheTestHandle>>() {
CacheBinding::Native(handle.service()?)
} else if let Some(service) = facade::resolve(py, &object)? {
CacheBinding::Native(service)
} else {
CacheBinding::PythonCallback(PythonCallback::new(object.unbind()))
};
Ok(ResolvedCache::new(binding))
ResolvedCache::from_selected(&object)
}
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {

View file

@ -8,16 +8,12 @@ mod logger;
mod marshal;
mod python_settings;
mod routes;
#[allow(
dead_code,
reason = "secret-manager foundations await rollout activation"
)]
mod secrets;
mod tokenizer;
#[pymodule(gil_used = true)]
mod _native {
use crate::cache::{CacheTestHandle, CacheTestResolver, ResolvedCache};
use crate::cache::{CacheResolver, CacheTestHandle, ResolvedCache};
#[cfg(feature = "panic-test")]
#[pymodule_export]
use crate::diagnostics::_panic_for_test;
@ -31,14 +27,16 @@ mod _native {
use crate::routes::audio_transcription::{atranscription, transcription};
#[pymodule_export]
use crate::routes::chat_completions::{
achat_completions, chat_completions, chat_completions_decline,
achat_completions, acompletion, chat_completions, chat_completions_decline, completion,
};
#[pymodule_export]
use crate::routes::embeddings::{aembedding, embedding};
#[pymodule_export]
use crate::routes::messages::{amessages, messages};
#[pymodule_export]
use crate::routes::ocr::{aocr, ocr};
#[pymodule_export]
use crate::routes::responses::ResponsesWebSocketConnection;
use crate::routes::responses::{ResponsesWebSocketConnection, aresponses, responses};
#[pymodule_export]
use crate::routes::token_counter::TokenCounter;
#[cfg(feature = "huggingface")]
@ -55,8 +53,13 @@ mod _native {
let py = module.py();
let dict = module.dict();
dict.set_item("_CacheTestHandle", py.get_type::<CacheTestHandle>())?;
dict.set_item("_CacheTestResolver", py.get_type::<CacheTestResolver>())?;
dict.set_item("_ResponseCacheRuntime", py.get_type::<ResolvedCache>())
dict.set_item("_CacheResolver", py.get_type::<CacheResolver>())?;
dict.set_item("_CacheTestResolver", py.get_type::<CacheResolver>())?;
dict.set_item("_ResponseCacheRuntime", py.get_type::<ResolvedCache>())?;
dict.set_item(
"_SecretManagerRuntime",
py.get_type::<crate::secrets::runtime::NativeSecretManager>(),
)
}
}
@ -82,6 +85,8 @@ mod tests {
"ProcessReservedForForking",
"ocr",
"aocr",
"embedding",
"aembedding",
"transcription",
"atranscription",
"messages",
@ -89,6 +94,10 @@ mod tests {
"chat_completions_decline",
"chat_completions",
"achat_completions",
"completion",
"acompletion",
"responses",
"aresponses",
"ResponsesWebSocketConnection",
"NativeDiagnosticProcessor",
"TokenCounter",

View file

@ -1,3 +1,6 @@
use pyo3::types::{PyDict, PyTuple};
use crate::errors::RustBridgeDeclined;
use crate::logger::{run_async, run_sync};
use litellm_core::chat_completions::{
Error, chat_completions as run_chat_completions, chat_completions_decline_reason,
@ -123,9 +126,58 @@ pub(crate) fn achat_completions<'py>(
)
}
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn completion(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native chat completions route is not implemented",
))
}
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn acompletion(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native chat completions route is not implemented",
))
}
#[cfg(test)]
mod tests {
use pyo3::{prelude::*, types::PyList};
use pyo3::{
prelude::*,
types::{PyDict, PyList, PyTuple},
};
use crate::errors::RustBridgeDeclined;
#[test]
fn both_entrypoints_decline_before_provider_execution() {
Python::initialize();
Python::attach(|py| {
let request = PyDict::new(py);
let args = PyTuple::empty(py);
let kwargs = PyDict::new(py);
for entrypoint in [super::completion, super::acompletion] {
let error = entrypoint(request.clone().into_any(), args.clone(), kwargs.clone())
.expect_err(
"native chat completions must decline until a route machine exists",
);
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
}
});
}
#[test]
fn chat_completions_decline_keeps_existing_reasons() {

View file

@ -0,0 +1,58 @@
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use crate::errors::RustBridgeDeclined;
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn embedding(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native embeddings route is not implemented",
))
}
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn aembedding(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native embeddings route is not implemented",
))
}
#[cfg(test)]
mod tests {
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use crate::errors::RustBridgeDeclined;
#[test]
fn both_entrypoints_decline_before_provider_execution() {
Python::initialize();
Python::attach(|py| {
let request = PyDict::new(py);
let args = PyTuple::empty(py);
let kwargs = PyDict::new(py);
for entrypoint in [super::embedding, super::aembedding] {
let error = entrypoint(request.clone().into_any(), args.clone(), kwargs.clone())
.expect_err("native embeddings must decline until a route machine exists");
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
}
});
}
}

View file

@ -1,5 +1,6 @@
pub(crate) mod audio_transcription;
pub(crate) mod chat_completions;
pub(crate) mod embeddings;
pub(crate) mod messages;
pub(crate) mod ocr;
pub(crate) mod responses;

View file

@ -3,17 +3,14 @@ mod errors;
mod host;
mod project;
use std::sync::{Arc, LazyLock};
use std::sync::LazyLock;
use host::OcrRouteHost;
use litellm_auth_gcp::VertexAuth;
use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call};
use litellm_core::ocr::route::ocr_machine;
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_llms::base_llm::{
inference::secrets::{EnvironmentSecrets, SecretSource},
ocr::{handler::OcrClient, settings::OcrSettings},
};
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
@ -21,14 +18,11 @@ use pyo3::{
use crate::{
coercion::FieldSpec,
errors::RustBridgeDeclined,
http,
python_settings::{PythonSettings, Snapshot},
secrets,
};
const SECRET_MANAGER_READABLE: FieldSpec<bool> =
FieldSpec::new("readable", |field| field.schema_bool());
const VERTEX_PROJECT: FieldSpec<Option<String>> =
FieldSpec::new("vertex_project", |field| field.falsy_optional_string());
const VERTEX_LOCATION: FieldSpec<Option<String>> =
@ -58,7 +52,7 @@ fn run_ocr(
kwargs: Bound<'_, PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let secrets = process_environment_secrets(&PythonSettings::SecretManager.read(py)?)?;
let secrets = secrets::source(py)?;
let config = http::call_config(py, &kwargs, asynchronous)?;
let client = OcrClient::new(
http::pool(),
@ -79,15 +73,6 @@ fn run_ocr(
)
}
fn process_environment_secrets(snapshot: &Snapshot<'_>) -> PyResult<Arc<dyn SecretSource>> {
if snapshot.read(&SECRET_MANAGER_READABLE)? {
return Err(RustBridgeDeclined::new_err(
"a readable secret manager is configured and the Rust route only reads the process environment",
));
}
Ok(Arc::new(EnvironmentSecrets))
}
fn ocr_settings(py: Python<'_>) -> PyResult<OcrSettings> {
project_provider_defaults(&PythonSettings::ProviderDefaults.read(py)?)
}
@ -123,38 +108,10 @@ pub(crate) fn aocr(
#[cfg(test)]
mod tests {
use pyo3::{prelude::*, types::PyDict};
use super::process_environment_secrets;
use crate::errors::RustBridgeDeclined;
use pyo3::prelude::*;
use crate::python_settings::PythonSettings;
fn secret_manager<'py>(py: Python<'py>, readable: bool) -> Bound<'py, PyAny> {
let locals = PyDict::new(py);
locals.set_item("readable", readable).unwrap();
py.run(
c"import types\nmanager = types.SimpleNamespace(readable=readable)",
Some(&locals),
Some(&locals),
)
.unwrap();
locals.get_item("manager").unwrap().unwrap()
}
#[test]
fn a_readable_secret_manager_sends_the_call_back_to_python() {
Python::initialize();
Python::attach(|py| {
let declined = process_environment_secrets(
&PythonSettings::SecretManager.snapshot(secret_manager(py, true)),
)
.err()
.expect("the Rust route declines");
assert!(declined.is_instance_of::<RustBridgeDeclined>(py));
});
}
#[test]
fn provider_defaults_distinguish_falsey_values_and_exact_true() {
Python::initialize();

View file

@ -1,12 +1,41 @@
use litellm_core::responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use pyo3::prelude::*;
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use serde_json::Value;
use crate::{
errors::responses_error_to_pyerr,
errors::{RustBridgeDeclined, responses_error_to_pyerr},
marshal::{marshal_headers, optional_timeout},
};
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn responses(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native responses route is not implemented",
))
}
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn aresponses(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native responses route is not implemented",
))
}
#[pyclass]
pub(crate) struct ResponsesWebSocketConnection {
inner: RustResponsesWebSocketConnection,
@ -63,7 +92,28 @@ mod tests {
use std::{ffi::CString, time::Duration};
use futures_util::{SinkExt, StreamExt};
use pyo3::{prelude::*, types::PyDict};
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use crate::errors::RustBridgeDeclined;
#[test]
fn both_entrypoints_decline_before_provider_execution() {
Python::initialize();
Python::attach(|py| {
let request = PyDict::new(py);
let args = PyTuple::empty(py);
let kwargs = PyDict::new(py);
for entrypoint in [super::responses, super::aresponses] {
let error = entrypoint(request.clone().into_any(), args.clone(), kwargs.clone())
.expect_err("native responses must decline until a route machine exists");
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
}
});
}
use tokio::net::TcpListener;
use tokio_tungstenite::{accept_async, tungstenite::Message};

View file

@ -4,11 +4,17 @@ use litellm_core_utils::settings::Lookup;
use litellm_secrets::{
Error, ExternalSecretManager, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue,
};
use pyo3::{prelude::*, types::PyDict};
use pyo3::{
exceptions::PyException,
prelude::*,
types::{PyDict, PyString},
};
use super::error::external_error;
use super::error::{external_error, read_error};
const HANDLER_MODULE: &str = "litellm.secret_managers.secret_manager_handler";
const ENVIRONMENT_FALLBACK_LOG: &str =
"Defaulting to os.environ value for key=%s. An exception occurred - %s.\n\n%s";
/// A secret manager whose reads execute in Python: a custom manager, a legacy compatible
/// client, or a manually assigned SDK client.
@ -33,23 +39,6 @@ impl PythonSecretManager {
fn read(&self, py: Python<'_>, name: &str) -> PyResult<Option<String>> {
let client = self.client.bind(py);
if self.system == Some(KeyManagementSystem::Custom)
|| (self.system.is_none() && client.hasattr("sync_read_secret")?)
{
let kwargs = PyDict::new(py);
kwargs.set_item("secret_name", name)?;
if self.system == Some(KeyManagementSystem::Custom) {
let optional_params = self
.settings
.as_ref()
.map(|settings| settings.bind(py).call_method0("model_dump"))
.transpose()?;
kwargs.set_item("optional_params", optional_params)?;
}
return client
.call_method("sync_read_secret", (), Some(&kwargs))?
.extract();
}
let kwargs = PyDict::new(py);
kwargs.set_item("client", client)?;
kwargs.set_item("key_manager", self.system.map_or("local", python_name))?;
@ -58,10 +47,15 @@ impl PythonSecretManager {
Some(settings) => kwargs.set_item("key_management_settings", settings.bind(py))?,
None => kwargs.set_item("key_management_settings", py.None())?,
}
py.import(HANDLER_MODULE)?
let result = py
.import(HANDLER_MODULE)?
.getattr("get_secret_from_manager")?
.call((), Some(&kwargs))?
.extract()
.call((), Some(&kwargs))?;
if result.is_instance_of::<PyString>() {
result.extract().map(Some)
} else {
Ok(None)
}
}
}
@ -92,15 +86,35 @@ impl ExternalSecretManager for PythonSecretManager {
_environment: &'a (dyn Lookup + Send + Sync),
) -> Pin<Box<dyn Future<Output = Result<Option<Secret>, Error>> + Send + 'a>> {
Box::pin(async move {
Python::attach(|py| {
self.read(py, name)
.map(|value| value.map(SecretValue::new).map(Secret::String))
.map_err(|error| external_error(py, error))
Python::attach(|py| match self.read(py, name) {
Ok(value) => Ok(value.map(SecretValue::new).map(Secret::String)),
// `get_secret` answers a failed manager read from the process environment, but
// only for `Exception`: cancellation and other `BaseException`s propagate.
Err(error) if error.is_instance_of::<PyException>(py) => {
log_environment_fallback(py, name, &error)
.map_err(|error| external_error(py, error))?;
Err(read_error(py, error))
}
Err(error) => Err(external_error(py, error)),
})
})
}
}
fn log_environment_fallback(py: Python<'_>, name: &str, error: &PyErr) -> PyResult<()> {
let traceback = py
.import("traceback")?
.call_method1("format_exception", (error.value(py),))?;
let traceback = "".into_pyobject(py)?.call_method1("join", (traceback,))?;
py.import("litellm._logging")?
.getattr("verbose_logger")?
.call_method1(
"error",
(ENVIRONMENT_FALLBACK_LOG, name, error.value(py), traceback),
)?;
Ok(())
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
@ -115,16 +129,12 @@ mod tests {
use super::{HANDLER_MODULE, PythonSecretManager, python_name};
use crate::secrets::python_error;
#[rstest]
#[case::value_error("ValueError", None)]
#[case::value_error_with_fallback("ValueError", Some("environment-key"))]
#[case::cancelled("asyncio.CancelledError", None)]
#[case::cancelled_with_fallback("asyncio.CancelledError", Some("environment-key"))]
#[tokio::test]
async fn callback_failures_preserve_python_exceptions_even_with_environment_fallback(
#[case] failure_type: &str,
#[case] fallback: Option<&'static str>,
) {
/// A resolver over a Python manager whose reads raise `failure_type`, with the chained
/// exceptions Python attaches, and `fallback` as the process environment.
fn failing_resolver(
failure_type: &str,
fallback: Option<&'static str>,
) -> (SecretResolver, Py<PyDict>) {
Python::initialize();
let (reader, locals) = Python::attach(|py| {
let locals = PyDict::new(py);
@ -141,6 +151,13 @@ class Manager:
def sync_read_secret(self, secret_name):
raise failure
manager = Manager()
import sys, types
for name in ('litellm', 'litellm.secret_managers'):
sys.modules.setdefault(name, types.ModuleType(name))
handler = sys.modules.setdefault('litellm.secret_managers.secret_manager_handler', types.ModuleType('litellm.secret_managers.secret_manager_handler'))
def get_secret_from_manager(**kwargs):
return kwargs['client'].sync_read_secret(kwargs['secret_name'])
handler.get_secret_from_manager = get_secret_from_manager
",
Some(&locals),
Some(&locals),
@ -153,7 +170,7 @@ manager = Manager()
);
(reader, locals.unbind())
});
let resolver = SecretResolver::new(
let resolver = SecretResolver::new_python_compatible(
Arc::new(SecretManagerState::new(
SecretManager::External(Arc::new(reader)),
KeyManagementSettings::default(),
@ -162,6 +179,19 @@ manager = Manager()
OidcResolver::default(),
)
.with_failure_policy(FailurePolicy::EnvironmentFallback);
(resolver, locals)
}
#[rstest]
#[case::cancelled("asyncio.CancelledError", None)]
#[case::cancelled_with_fallback("asyncio.CancelledError", Some("environment-key"))]
#[case::keyboard_interrupt("KeyboardInterrupt", Some("environment-key"))]
#[tokio::test]
async fn base_exceptions_propagate_unchanged_even_with_environment_fallback(
#[case] failure_type: &str,
#[case] fallback: Option<&'static str>,
) {
let (resolver, locals) = failing_resolver(failure_type, fallback);
let error = resolver.get_secret("API_KEY", None).await.unwrap_err();
Python::attach(|py| {
let original = python_error(py, &error).unwrap();
@ -184,24 +214,96 @@ manager = Manager()
});
}
/// Installs a persistent `litellm._logging` stub whose `verbose_logger.error` records its
/// arguments, and returns those recorded for `name`.
fn logged_errors<'py>(py: Python<'py>, name: &str) -> Vec<Bound<'py, PyAny>> {
py.run(
c"
import sys, types
class Logger:
calls = []
def error(self, *args):
self.calls.append(args)
logging = types.ModuleType('litellm._logging')
logging.verbose_logger = Logger()
sys.modules.setdefault('litellm', types.ModuleType('litellm'))
sys.modules.setdefault('litellm._logging', logging)
",
None,
None,
)
.unwrap();
py.import("litellm._logging")
.unwrap()
.getattr("verbose_logger")
.unwrap()
.getattr("calls")
.unwrap()
.try_iter()
.unwrap()
.map(Result::unwrap)
.filter(|call| call.get_item(1).unwrap().extract::<String>().unwrap() == name)
.collect()
}
#[rstest]
#[case::value_error("ValueError", None, "FALLBACK_VALUE_ERROR")]
#[case::value_error_with_fallback(
"ValueError",
Some("environment-key"),
"FALLBACK_VALUE_ERROR_WITH_ENVIRONMENT"
)]
#[case::runtime_error_with_fallback(
"RuntimeError",
Some("environment-key"),
"FALLBACK_RUNTIME_ERROR_WITH_ENVIRONMENT"
)]
#[tokio::test]
async fn exceptions_are_logged_and_answered_from_the_environment(
#[case] failure_type: &str,
#[case] fallback: Option<&'static str>,
#[case] name: &str,
) {
let (resolver, _locals) = failing_resolver(failure_type, fallback);
Python::attach(|py| assert!(logged_errors(py, name).is_empty()));
let secret = resolver.get_secret(name, None).await.unwrap();
assert_eq!(
secret.map(|secret| match secret {
litellm_secrets::Secret::String(value) => value.expose().to_owned(),
other => panic!("unexpected secret {other:?}"),
}),
fallback.map(str::to_owned)
);
Python::attach(|py| {
let calls = logged_errors(py, name);
assert_eq!(calls.len(), 1);
assert!(
calls[0]
.get_item(3)
.unwrap()
.extract::<String>()
.unwrap()
.contains("sync_read_secret")
);
});
}
/// Installs a fake `get_secret_from_manager` that records its kwargs, runs `body`, and
/// removes the fake modules again.
/// removes the fake handler again; parent package stubs persist for concurrent tests.
fn with_fake_handler<'py>(py: Python<'py>, body: impl FnOnce(&Bound<'py, PyDict>)) {
let locals = PyDict::new(py);
py.run(
c"
import sys, types
previous_handler = sys.modules.get('litellm.secret_managers.secret_manager_handler')
calls = []
def get_secret_from_manager(**kwargs):
calls.append(kwargs)
return 'handled-' + kwargs['secret_name']
handler = types.ModuleType('litellm.secret_managers.secret_manager_handler')
handler.get_secret_from_manager = get_secret_from_manager
installed = {}
for name in ('litellm', 'litellm.secret_managers'):
if name not in sys.modules:
sys.modules[name] = types.ModuleType(name)
installed[name] = True
sys.modules.setdefault(name, types.ModuleType(name))
sys.modules['litellm.secret_managers.secret_manager_handler'] = handler
",
Some(&locals),
@ -211,9 +313,10 @@ sys.modules['litellm.secret_managers.secret_manager_handler'] = handler
body(&locals);
py.run(
c"
sys.modules.pop('litellm.secret_managers.secret_manager_handler', None)
for name in installed:
sys.modules.pop(name, None)
if previous_handler is None:
sys.modules.pop('litellm.secret_managers.secret_manager_handler', None)
else:
sys.modules['litellm.secret_managers.secret_manager_handler'] = previous_handler
",
Some(&locals),
Some(&locals),
@ -221,6 +324,28 @@ for name in installed:
.unwrap();
}
#[rstest]
#[case("None")]
#[case("True")]
#[case("123")]
#[case("{'key': 'value'}")]
fn nonstring_results_are_absent_without_a_read_failure(#[case] expression: &str) {
Python::initialize();
Python::attach(|py| {
with_fake_handler(py, |locals| {
locals.set_item("expression", expression).unwrap();
py.run(
c"handler.get_secret_from_manager = lambda **kwargs: eval(expression)",
Some(locals),
Some(locals),
)
.unwrap();
let reader = PythonSecretManager::new(py.None(), None, None);
assert_eq!(reader.read(py, "KEY").unwrap(), None);
});
});
}
#[rstest]
#[case::google_kms(KeyManagementSystem::GoogleKms)]
#[case::azure_key_vault(KeyManagementSystem::AzureKeyVault)]
@ -238,72 +363,6 @@ for name in installed:
);
}
#[rstest]
#[case::legacy(None, false)]
#[case::custom(Some(KeyManagementSystem::Custom), true)]
fn direct_readers_receive_compatible_kwargs(
#[case] system: Option<KeyManagementSystem>,
#[case] expects_optional_params: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"
class Settings:
def model_dump(self):
return {'scope': 'custom'}
class Manager:
def __init__(self):
self.names = []
self.optional_params = []
def sync_read_secret(self, secret_name, optional_params=None, timeout=None):
self.names.append(secret_name)
self.optional_params.append(optional_params)
return 'direct-' + secret_name
manager = Manager()
settings = Settings()
",
Some(&locals),
Some(&locals),
)
.unwrap();
let manager = locals.get_item("manager").unwrap().unwrap();
let settings = expects_optional_params
.then(|| locals.get_item("settings").unwrap().unwrap().unbind());
let reader = PythonSecretManager::new(manager.clone().unbind(), system, settings);
assert_eq!(
reader.read(py, "API_KEY").unwrap().as_deref(),
Some("direct-API_KEY")
);
assert_eq!(
manager
.getattr("names")
.unwrap()
.extract::<Vec<String>>()
.unwrap(),
["API_KEY"]
);
let optional_params = manager
.getattr("optional_params")
.unwrap()
.get_item(0)
.unwrap();
if expects_optional_params {
assert_eq!(
optional_params
.get_item("scope")
.unwrap()
.extract::<String>()
.unwrap(),
"custom"
);
} else {
assert!(optional_params.is_none());
}
});
}
#[test]
fn configured_systems_dispatch_through_the_python_handler_with_the_original_settings() {
Python::initialize();
@ -349,4 +408,57 @@ settings = Settings()
});
});
}
#[rstest]
#[case::manually_assigned(None, "local")]
#[case::custom(Some(KeyManagementSystem::Custom), "custom")]
fn direct_readers_dispatch_through_the_python_handler_like_get_secret(
#[case] system: Option<KeyManagementSystem>,
#[case] key_manager: &str,
) {
Python::initialize();
Python::attach(|py| {
with_fake_handler(py, |locals| {
py.run(
c"
class Manager:
def __init__(self):
self.names = []
def sync_read_secret(self, secret_name, optional_params=None, timeout=None):
self.names.append(secret_name)
return 'direct-' + secret_name
manager = Manager()
",
Some(locals),
Some(locals),
)
.unwrap();
let manager = locals.get_item("manager").unwrap().unwrap();
let reader = PythonSecretManager::new(manager.clone().unbind(), system, None);
assert_eq!(
reader.read(py, "API_KEY").unwrap().as_deref(),
Some("handled-API_KEY")
);
assert_eq!(
manager
.getattr("names")
.unwrap()
.extract::<Vec<String>>()
.unwrap(),
Vec::<String>::new()
);
let calls = locals.get_item("calls").unwrap().unwrap();
let call = calls.get_item(0).unwrap().cast_into::<PyDict>().unwrap();
assert!(call.get_item("client").unwrap().unwrap().is(&manager));
assert_eq!(
call.get_item("key_manager")
.unwrap()
.unwrap()
.extract::<String>()
.unwrap(),
key_manager
);
});
});
}
}

View file

@ -56,13 +56,23 @@ const SETTINGS_OBJECT: FieldSpec<Option<Py<PyAny>>> =
FieldSpec::new("settings_object", |field| Ok(field.python_binding()));
/// `litellm.secret_manager_client` as the bridge classifies it.
#[derive(Debug)]
pub(crate) enum SecretManagerClient {
/// `None`: reads come from the process environment.
Local,
/// A custom manager, legacy compatible client, or manually assigned SDK client that keeps
/// executing in Python.
PythonCallback(Py<PyAny>),
Native(Box<SecretManager>),
}
impl std::fmt::Debug for SecretManagerClient {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str(match self {
Self::Local => "Local",
Self::Native(_) => "Native",
Self::PythonCallback(_) => "PythonCallback",
})
}
}
/// One operation-local capture of the secret manager globals, taken while attached to Python.
@ -79,6 +89,9 @@ pub(crate) struct SecretManagerSnapshot {
impl SecretManagerSnapshot {
pub(crate) fn into_state(self) -> Arc<SecretManagerState> {
match self.client {
SecretManagerClient::Native(backend) => {
Arc::new(SecretManagerState::new(*backend, self.settings))
}
SecretManagerClient::Local => Arc::new(SecretManagerState::default()),
SecretManagerClient::PythonCallback(client) => Arc::new(SecretManagerState::new(
SecretManager::External(Arc::new(PythonSecretManager::new(
@ -94,7 +107,32 @@ impl SecretManagerSnapshot {
/// Reads and projects the secret manager settings group in one attached operation.
pub(crate) fn read(py: Python<'_>) -> PyResult<SecretManagerSnapshot> {
Ok(project(&PythonSettings::SecretManagerBinding.read(py)?)?)
let snapshot = project(&PythonSettings::SecretManagerBinding.read(py)?)?;
let SecretManagerClient::PythonCallback(client) = &snapshot.client else {
return Ok(snapshot);
};
if matches!(
snapshot.system,
Some(KeyManagementSystem::Custom | KeyManagementSystem::Local)
) {
return Ok(snapshot);
}
let Some(native) = super::runtime::NativeSecretManager::from_client(client.bind(py))? else {
return Ok(snapshot);
};
let backend = native.borrow(py).backend()?;
if snapshot
.system
.is_some_and(|system| system != backend.system())
{
return Err(pyo3::exceptions::PyValueError::new_err(
"native secret manager system does not match configuration",
));
}
Ok(SecretManagerSnapshot {
client: SecretManagerClient::Native(Box::new(backend)),
..snapshot
})
}
pub(crate) fn project(snapshot: &Snapshot<'_>) -> Result<SecretManagerSnapshot, ProjectionError> {

View file

@ -17,8 +17,12 @@ pub(super) fn external_error(py: Python<'_>, error: PyErr) -> Error {
Error::ExternalManager(Box::new(PythonSecretError(error.into_value(py))))
}
pub(super) fn read_error(py: Python<'_>, error: PyErr) -> Error {
Error::ExternalRead(Box::new(PythonSecretError(error.into_value(py))))
}
pub(crate) fn python_error(py: Python<'_>, error: &Error) -> Option<PyErr> {
let Error::ExternalManager(source) = error else {
let (Error::ExternalManager(source) | Error::ExternalRead(source)) = error else {
return None;
};
source

View file

@ -1,6 +1,105 @@
pub(crate) mod callback;
pub(crate) mod config;
mod error;
mod mutation;
mod operations;
mod provider;
pub(crate) mod resolved;
pub(crate) mod runtime;
mod vault;
use std::sync::Arc;
use litellm_secrets::source::{EnvironmentSecrets, SecretSource};
use pyo3::prelude::*;
pub(crate) use error::python_error;
use resolved::ResolvedSecrets;
use crate::{
coercion::FieldSpec,
errors::RustBridgeDeclined,
python_settings::{PythonSettings, Snapshot},
};
const READABLE: FieldSpec<bool> = FieldSpec::new("readable", |field| field.schema_bool());
const NATIVE: FieldSpec<bool> = FieldSpec::new("native", |field| field.schema_bool());
/// Where a Rust route reads provider secrets from, as `litellm.get_secret` would.
pub(crate) fn source(py: Python<'_>) -> PyResult<Arc<dyn SecretSource>> {
select(&PythonSettings::SecretManager.read(py)?, || {
Ok(Arc::new(ResolvedSecrets::new(config::read(py)?)))
})
}
fn select(
manager: &Snapshot<'_>,
resolved: impl FnOnce() -> PyResult<Arc<dyn SecretSource>>,
) -> PyResult<Arc<dyn SecretSource>> {
if !manager.read(&READABLE)? {
return Ok(Arc::new(EnvironmentSecrets::python_compatible()));
}
if !manager.read(&NATIVE)? {
return Err(RustBridgeDeclined::new_err(
"the configured secret manager is not enabled for the Rust bridge",
));
}
resolved()
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use litellm_secrets::source::{EnvironmentSecrets, SecretSource};
use pyo3::{prelude::*, types::PyDict};
use rstest::rstest;
use super::select;
use crate::{errors::RustBridgeDeclined, python_settings::PythonSettings};
enum Selected {
Environment,
Declined,
Resolved,
}
#[rstest]
#[case::unreadable(false, false, Selected::Environment)]
#[case::unreadable_even_if_native(false, true, Selected::Environment)]
#[case::readable_python_only(true, false, Selected::Declined)]
#[case::readable_native(true, true, Selected::Resolved)]
fn readable_and_native_select_the_secret_source(
#[case] readable: bool,
#[case] native: bool,
#[case] expected: Selected,
) {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
locals.set_item("readable", readable).unwrap();
locals.set_item("native", native).unwrap();
let manager = py
.eval(
c"__import__('types').SimpleNamespace(readable=readable, native=native)",
None,
Some(&locals),
)
.unwrap();
let mut resolved_called = false;
let selected = select(&PythonSettings::SecretManager.snapshot(manager), || {
resolved_called = true;
Ok(Arc::new(EnvironmentSecrets::python_compatible()) as Arc<dyn SecretSource>)
});
match expected {
Selected::Environment => assert!(selected.is_ok() && !resolved_called),
Selected::Resolved => assert!(selected.is_ok() && resolved_called),
Selected::Declined => {
let error = selected.err().expect("the Rust route declines");
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
assert!(!resolved_called);
}
}
});
}
}

View file

@ -0,0 +1,113 @@
use super::operations::{PythonMutationError, PythonMutationResponse};
use litellm_host_python::{json_loads, to_py};
use litellm_secrets::cyberark;
use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict};
pub(super) fn mutation_value(
result: Result<PythonMutationResponse, PythonMutationError>,
context: &super::vault::ErrorContext,
) -> PyResult<Py<PyAny>> {
Python::attach(|py| match result {
Ok(PythonMutationResponse::Value(value)) => to_py(py, &value),
Ok(PythonMutationResponse::Json(body)) => match json_value(py, &body) {
Ok(value) => Ok(value),
Err(error) => error_value(py, error.value(py).str()?.extract()?),
},
Err(PythonMutationError::Vault(failure)) => {
super::vault::failure_value(py, *failure, context)
}
Err(PythonMutationError::CyberarkWrite { name, failure }) => {
let message = cyberark_failure(py, &name, *failure)?;
to_py(
py,
&serde_json::json!({"status": "error", "message": message}),
)
}
Err(PythonMutationError::CurrentMissing(name)) => Err(PyValueError::new_err(format!(
"Current secret {name} not found"
))),
Err(PythonMutationError::ReplacementMissing(name)) => Err(PyValueError::new_err(format!(
"Failed to verify new secret {name}"
))),
Err(PythonMutationError::ReplacementMismatch) => {
Err(PyValueError::new_err("New secret value mismatch"))
}
Err(PythonMutationError::Unsupported) => Err(PyValueError::new_err(
"native secret manager mutation is unavailable",
)),
})
}
fn cyberark_failure(
py: Python<'_>,
name: &str,
failure: cyberark::WriteFailure,
) -> PyResult<String> {
let message = match failure.source {
cyberark::Error::Status(status) | cyberark::Error::AuthStatus(status) => {
let url = failure
.request_url
.as_ref()
.map_or("", reqwest::Url::as_str);
http_message(py, "POST", url, status)?
}
cyberark::Error::Operation(litellm_secrets_types::Error::UnsafeSecretName) => {
format!("Invalid secret_name {}", name.into_pyobject(py)?.repr()?)
}
cyberark::Error::Http(source) if failure.authentication => match os_error_code(&source) {
Some(code) => {
let reason = py.import("os")?.getattr("strerror")?.call1((code,))?;
py.import("builtins")?
.getattr("OSError")?
.call1((code, reason))?
.str()?
.extract()?
}
None => cyberark::Error::Http(source).to_string(),
},
cyberark::Error::Http(source) if source.is_connect() => {
"All connection attempts failed".to_owned()
}
source => source.to_string(),
};
Ok(if failure.authentication {
format!("Could not authenticate to CyberArk Conjur: {message}")
} else {
message
})
}
fn os_error_code(error: &(dyn std::error::Error + 'static)) -> Option<i32> {
error
.downcast_ref::<std::io::Error>()
.and_then(std::io::Error::raw_os_error)
.or_else(|| error.source().and_then(os_error_code))
}
pub(super) fn json_value(py: Python<'_>, body: &[u8]) -> PyResult<Py<PyAny>> {
json_loads(py, body)
}
pub(super) fn error_value(py: Python<'_>, message: String) -> PyResult<Py<PyAny>> {
to_py(
py,
&serde_json::json!({"status": "error", "message": message}),
)
}
pub(super) fn http_message(
py: Python<'_>,
method: &str,
url: &str,
status: u16,
) -> PyResult<String> {
let httpx = py.import("httpx")?;
let request = httpx.getattr("Request")?.call1((method, url))?;
let kwargs = PyDict::new(py);
kwargs.set_item("request", request)?;
let response = httpx.getattr("Response")?.call((status,), Some(&kwargs))?;
match response.call_method0("raise_for_status") {
Err(error) => error.value(py).str()?.extract(),
Ok(_) => Ok(format!("HTTP {status}")),
}
}

View file

@ -0,0 +1,247 @@
use litellm_core_utils::settings::Lookup;
use litellm_secrets::Secret;
use litellm_secrets::cyberark::AuthenticationRetry;
use litellm_secrets::{Error, SecretManager};
use litellm_secrets_types::{PythonSecretRead, SecretOperationContext};
pub(super) struct PythonReadRequest {
pub secret_name: String,
pub primary_secret_name: Option<String>,
pub context: SecretOperationContext,
pub synchronous: bool,
}
pub(super) async fn read_python_provider(
manager: &SecretManager,
request: &PythonReadRequest,
_environment: &(dyn Lookup + Send + Sync),
) -> Result<PythonSecretRead, Error> {
match (manager, &request.context) {
(SecretManager::AwsSecretsManagerV2(client), SecretOperationContext::Aws(context)) => {
client
.read_provider_payload_for_python(
&request.secret_name,
request.primary_secret_name.as_deref(),
context,
request.synchronous,
_environment,
)
.await
.map_err(Error::from)
}
(SecretManager::HashicorpVault(client), SecretOperationContext::Hashicorp(context)) => {
Ok(PythonSecretRead::Value(
client
.async_read_secret_with_context(&request.secret_name, context)
.await
.unwrap_or(None)
.map(Secret::String),
))
}
(SecretManager::Cyberark(client), SecretOperationContext::Cyberark(_)) => {
Ok(PythonSecretRead::Value(
client
.read_with_retry(
&request.secret_name,
&Default::default(),
AuthenticationRetry::Never,
)
.await
.unwrap_or(None)
.map(Secret::String),
))
}
(SecretManager::GoogleSecretManager(client), SecretOperationContext::Google(_)) => client
.get_secret_for_python(&request.secret_name)
.await
.map(PythonSecretRead::Value)
.map_err(Error::from),
_ => Err(Error::NativeBackendUnavailable),
}
}
#[derive(Debug)]
pub(super) enum PythonMutationError {
Unsupported,
Vault(Box<super::vault::Failure>),
CyberarkWrite {
name: String,
failure: Box<litellm_secrets::cyberark::WriteFailure>,
},
CurrentMissing(String),
ReplacementMissing(String),
ReplacementMismatch,
}
pub(super) async fn write_python_provider(
manager: &SecretManager,
name: &str,
value: &litellm_secrets::SecretValue,
) -> Result<serde_json::Value, PythonMutationError> {
match manager {
SecretManager::Cyberark(client) => {
client
.write_with_retry(name, value, &Default::default(), AuthenticationRetry::Never)
.await
.map_err(|failure| PythonMutationError::CyberarkWrite {
name: name.to_owned(),
failure: Box::new(failure),
})?;
Ok(write_success(name))
}
_ => Err(PythonMutationError::Unsupported),
}
}
pub(super) async fn delete_python_provider(
manager: &SecretManager,
name: &str,
) -> Result<serde_json::Value, PythonMutationError> {
match manager {
SecretManager::Cyberark(client) => {
client
.async_delete_secret(name, None)
.await
.map_err(|failure| PythonMutationError::CyberarkWrite {
name: name.to_owned(),
failure: Box::new(litellm_secrets::cyberark::WriteFailure {
source: failure,
request_url: None,
authentication: false,
}),
})?;
Ok(serde_json::json!({
"status": "not_supported",
"message": "CyberArk Conjur does not support direct secret deletion. Use policy updates to remove variables.",
}))
}
_ => Err(PythonMutationError::Unsupported),
}
}
pub(super) async fn rotate_python_provider(
manager: &SecretManager,
current_name: &str,
new_name: &str,
value: &litellm_secrets::SecretValue,
) -> Result<serde_json::Value, PythonMutationError> {
match manager {
SecretManager::Cyberark(client) => {
if client
.read_fresh_with_retry(
current_name,
&Default::default(),
AuthenticationRetry::Never,
)
.await
.ok()
.flatten()
.is_none()
{
return Err(PythonMutationError::CurrentMissing(current_name.to_owned()));
}
client
.write_with_retry(
new_name,
value,
&Default::default(),
AuthenticationRetry::Never,
)
.await
.map_err(|failure| PythonMutationError::CyberarkWrite {
name: new_name.to_owned(),
failure: Box::new(failure),
})?;
let actual = client
.read_fresh_with_retry(new_name, &Default::default(), AuthenticationRetry::Never)
.await
.ok()
.flatten()
.ok_or_else(|| PythonMutationError::ReplacementMissing(new_name.to_owned()))?;
if actual != *value {
return Err(PythonMutationError::ReplacementMismatch);
}
if current_name != new_name {
client.invalidate_cached_secret(current_name).await;
}
Ok(write_success(new_name))
}
_ => Err(PythonMutationError::Unsupported),
}
}
fn write_success(name: &str) -> serde_json::Value {
serde_json::json!({"status": "success", "message": format!("Secret {name} written successfully")})
}
pub(super) enum PythonMutationResponse {
Value(serde_json::Value),
Json(Vec<u8>),
}
pub(super) async fn write_python_provider_with_context(
manager: &SecretManager,
name: &str,
value: &litellm_secrets::SecretValue,
context: &litellm_secrets_types::SecretWriteContext,
) -> Result<PythonMutationResponse, PythonMutationError> {
if let (SecretManager::HashicorpVault(client), SecretOperationContext::Hashicorp(operation)) =
(manager, &context.operation)
{
return super::vault::write(
client,
name,
value,
&litellm_secrets_types::SecretWriteContext {
description: context.description.clone(),
tags: context.tags.clone(),
operation: operation.clone(),
},
)
.await
.map(PythonMutationResponse::Json)
.map_err(|failure| PythonMutationError::Vault(Box::new(failure)));
}
write_python_provider(manager, name, value)
.await
.map(PythonMutationResponse::Value)
}
pub(super) async fn delete_python_provider_with_context(
manager: &SecretManager,
name: &str,
context: &SecretOperationContext,
) -> Result<PythonMutationResponse, PythonMutationError> {
if let (SecretManager::HashicorpVault(client), SecretOperationContext::Hashicorp(context)) =
(manager, context)
{
super::vault::delete(client, name, context)
.await
.map_err(|failure| PythonMutationError::Vault(Box::new(failure)))?;
return Ok(PythonMutationResponse::Value(serde_json::json!({
"status": "success", "message": format!("Secret {name} deleted successfully"),
})));
}
delete_python_provider(manager, name)
.await
.map(PythonMutationResponse::Value)
}
pub(super) async fn rotate_python_provider_with_context(
manager: &SecretManager,
current_name: &str,
new_name: &str,
value: &litellm_secrets::SecretValue,
context: &SecretOperationContext,
) -> Result<PythonMutationResponse, PythonMutationError> {
if let (SecretManager::HashicorpVault(client), SecretOperationContext::Hashicorp(context)) =
(manager, context)
{
return super::vault::rotate(client, current_name, new_name, value, context)
.await
.map(PythonMutationResponse::Json)
.map_err(|failure| PythonMutationError::Vault(Box::new(failure)));
}
rotate_python_provider(manager, current_name, new_name, value)
.await
.map(PythonMutationResponse::Value)
}

View file

@ -0,0 +1,150 @@
use std::time::Duration;
use super::operations::PythonReadRequest;
use litellm_secrets::{KeyManagementSystem, SecretValue};
use litellm_secrets_types::{
AwsOperationContext, CyberarkOperationContext, GoogleOperationContext,
HashicorpOperationContext, SecretOperationContext,
};
use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict};
pub(super) fn read_request(
system: KeyManagementSystem,
secret_name: String,
optional_params: Option<&Bound<'_, PyAny>>,
timeout: Option<&Bound<'_, PyAny>>,
primary_secret_name: Option<String>,
synchronous: bool,
) -> PyResult<PythonReadRequest> {
let context = match system {
KeyManagementSystem::AwsSecretManager => {
let ignored = primary_secret_name
.as_ref()
.is_some_and(|value| !value.is_empty())
|| (synchronous
&& litellm_secrets::aws::secret_manager::is_bootstrap_key(&secret_name));
SecretOperationContext::Aws(if ignored {
AwsOperationContext::default()
} else {
aws_context(optional_params, timeout)?
})
}
KeyManagementSystem::HashicorpVault => {
SecretOperationContext::Hashicorp(vault_context(optional_params)?)
}
KeyManagementSystem::Cyberark => {
SecretOperationContext::Cyberark(CyberarkOperationContext::default())
}
KeyManagementSystem::GoogleSecretManager => {
SecretOperationContext::Google(GoogleOperationContext::default())
}
_ => {
return Err(PyValueError::new_err(
"secret manager does not support provider reads",
));
}
};
Ok(PythonReadRequest {
secret_name,
primary_secret_name,
context,
synchronous,
})
}
fn string_field(params: Option<&Bound<'_, PyDict>>, name: &str) -> PyResult<Option<String>> {
let value = params
.map(|params| params.get_item(name))
.transpose()?
.flatten();
match value {
Some(value) if value.is_truthy()? => value.extract().map(Some),
_ => Ok(None),
}
}
fn aws_context(
params: Option<&Bound<'_, PyAny>>,
timeout: Option<&Bound<'_, PyAny>>,
) -> PyResult<AwsOperationContext> {
let params = params
.filter(|value| !value.is_none())
.map(|value| value.cast::<PyDict>())
.transpose()?;
Ok(AwsOperationContext {
access_key_id: string_field(params, "aws_access_key_id")?.map(SecretValue::new),
secret_access_key: string_field(params, "aws_secret_access_key")?.map(SecretValue::new),
session_token: string_field(params, "aws_session_token")?.map(SecretValue::new),
region_name: string_field(params, "aws_region_name")?,
role_name: string_field(params, "aws_role_name")?,
session_name: string_field(params, "aws_session_name")?,
external_id: string_field(params, "aws_external_id")?.map(SecretValue::new),
profile_name: string_field(params, "aws_profile_name")?,
web_identity_token: string_field(params, "aws_web_identity_token")?.map(SecretValue::new),
sts_endpoint: string_field(params, "aws_sts_endpoint")?,
bedrock_runtime_endpoint: string_field(params, "aws_bedrock_runtime_endpoint")?,
timeout: read_timeout(timeout)?,
})
}
fn read_timeout(value: Option<&Bound<'_, PyAny>>) -> PyResult<Option<Duration>> {
let Some(value) = value.filter(|value| !value.is_none()) else {
return Ok(None);
};
let seconds = match value.extract::<f64>() {
Ok(value) => Some(value),
Err(_) => value.getattr("read")?.extract::<Option<f64>>()?,
};
seconds
.map(|value| {
Duration::try_from_secs_f64(value).map_err(|_| PyValueError::new_err("invalid timeout"))
})
.transpose()
}
fn vault_context(params: Option<&Bound<'_, PyAny>>) -> PyResult<HashicorpOperationContext> {
let params = params.and_then(|value| value.cast::<PyDict>().ok());
let nested = params
.map(|params| params.get_item("secret_manager_settings"))
.transpose()?
.flatten();
let source = nested
.as_ref()
.and_then(|value| value.cast::<PyDict>().ok())
.or(params);
Ok(HashicorpOperationContext {
namespace: vault_field(source, "namespace")?,
mount: vault_field(source, "mount")?,
path_prefix: vault_field(source, "path_prefix")?,
data_key: vault_field(source, "data")?,
timeout: None,
})
}
fn vault_field(params: Option<&Bound<'_, PyDict>>, name: &str) -> PyResult<Option<String>> {
let value = params
.map(|params| params.get_item(name))
.transpose()?
.flatten();
match value {
Some(value) if value.is_none() => Ok(None),
Some(value) => value.str()?.extract().map(Some),
None => Ok(None),
}
}
pub(super) fn mutation_context(
system: KeyManagementSystem,
optional_params: Option<&Bound<'_, PyAny>>,
timeout: Option<&Bound<'_, PyAny>>,
) -> PyResult<SecretOperationContext> {
if system == KeyManagementSystem::HashicorpVault {
return Ok(SecretOperationContext::Hashicorp(
HashicorpOperationContext {
timeout: read_timeout(timeout)?,
..vault_context(optional_params)?
},
));
}
Ok(SecretOperationContext::Default)
}

View file

@ -1,10 +1,10 @@
use std::{collections::HashMap, sync::Arc};
use std::sync::Arc;
use futures_util::{future::BoxFuture, future::try_join_all};
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
use litellm_llms::base_llm::inference::secrets::{SecretSource, Secrets};
use futures_util::future::BoxFuture;
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_secrets::source::SecretSource;
use litellm_secrets::{
Error, FailurePolicy, OidcResolver, Secret, SecretManagerState, SecretResolver,
Error, FailurePolicy, OidcResolver, SecretManagerState, SecretResolver, SecretValue,
};
use super::config::SecretManagerSnapshot;
@ -20,7 +20,7 @@ impl ResolvedSecrets {
fn from_state(state: Arc<SecretManagerState>) -> Self {
Self {
resolver: SecretResolver::new(
resolver: SecretResolver::new_python_compatible(
state,
Arc::new(ProcessEnvironment),
OidcResolver::default(),
@ -31,41 +31,11 @@ impl ResolvedSecrets {
}
impl SecretSource for ResolvedSecrets {
fn resolve<'a>(&'a self, names: &'a [&'static str]) -> BoxFuture<'a, Result<Secrets, Error>> {
Box::pin(async move {
let values = try_join_all(names.iter().map(|name| async move {
self.resolver
.get_secret(name, None)
.await
.map(|secret| secret.map(|secret| ((*name).to_owned(), secret_value(secret))))
}))
.await?
.into_iter()
.flatten()
.collect::<HashMap<_, _>>();
Ok(Arc::new(ResolvedLookup { values }) as Secrets)
})
}
}
struct ResolvedLookup {
values: HashMap<String, String>,
}
impl Lookup for ResolvedLookup {
fn get(&self, name: &str) -> Option<String> {
self.values
.get(name)
.cloned()
.or_else(|| ProcessEnvironment.get(name))
}
}
fn secret_value(secret: Secret) -> String {
match secret {
Secret::String(value) => value.expose().to_owned(),
Secret::Bool(value) => if value { "True" } else { "False" }.to_owned(),
Secret::Json(value) => value.to_string(),
fn get_secret_str<'a>(
&'a self,
name: &'a str,
) -> BoxFuture<'a, Result<Option<SecretValue>, Error>> {
Box::pin(self.resolver.get_secret_str(name, None))
}
}
@ -86,7 +56,7 @@ mod tests {
};
use super::ResolvedSecrets;
use litellm_llms::base_llm::inference::secrets::SecretSource;
use litellm_secrets::source::SecretSource;
fn state(server: &MockServer, settings: KeyManagementSettings) -> Arc<SecretManagerState> {
let client = Client::from_conf(
@ -145,31 +115,75 @@ mod tests {
}
#[tokio::test]
async fn manager_failure_falls_back_to_environment() {
async fn aws_read_failure_preserves_absence_without_environment_fallback() {
let name = "LITELLM_RUST_BRIDGE_MANAGER_FAILURE";
unsafe { std::env::set_var(name, "env-key") };
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.respond_with(ResponseTemplate::new(500))
.expect(1)
.expect(2)
.mount(&server)
.await;
let result = resolve(state(&server, KeyManagementSettings::default()), name).await;
let missing = resolve(
state(&server, KeyManagementSettings::default()),
"LITELLM_RUST_BRIDGE_MANAGER_FAILURE_MISSING",
)
.await;
unsafe { std::env::remove_var(name) };
assert_eq!(result.as_deref(), Some("env-key"));
assert_eq!(server.received_requests().await.unwrap().len(), 1);
assert_eq!(result, None);
assert_eq!(missing, None);
}
let missing_server = MockServer::start().await;
#[rstest::rstest]
#[case::capitalized_true("True")]
#[case::parenthesized_false("(False)")]
#[tokio::test]
async fn boolean_manager_values_are_absent_like_get_secret_str(#[case] value: &str) {
let name = "LITELLM_RUST_BRIDGE_BOOLEAN_VALUE";
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.respond_with(ResponseTemplate::new(500))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString": value})))
.expect(1)
.mount(&missing_server)
.mount(&server)
.await;
let missing =
ResolvedSecrets::from_state(state(&missing_server, KeyManagementSettings::default()))
.resolve(&["LITELLM_RUST_BRIDGE_MANAGER_FAILURE_MISSING"])
.await;
assert!(matches!(missing, Err(litellm_secrets::Error::Aws(_))));
assert_eq!(
resolve(state(&server, KeyManagementSettings::default()), name).await,
None
);
}
#[tokio::test]
async fn undeclared_names_are_read_from_the_manager() {
let declared = "LITELLM_RUST_BRIDGE_DECLARED";
let undeclared = "LITELLM_RUST_BRIDGE_UNDECLARED_MANAGED";
unsafe { std::env::set_var(undeclared, "env-key") };
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.and(body_partial_json(json!({"SecretId": declared})))
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"SecretString": "declared-key"})),
)
.mount(&server)
.await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.and(body_partial_json(json!({"SecretId": undeclared})))
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"SecretString": "manager-key"})),
)
.expect(1)
.mount(&server)
.await;
let source = ResolvedSecrets::from_state(state(&server, KeyManagementSettings::default()));
let snapshot = source.resolve(&[declared]).await.unwrap();
assert_eq!(snapshot.get(undeclared), None);
let result = source
.get_secret_str(undeclared)
.await
.unwrap()
.map(|value| value.expose().to_owned());
unsafe { std::env::remove_var(undeclared) };
assert_eq!(result.as_deref(), Some("manager-key"));
}
#[tokio::test]
@ -230,7 +244,7 @@ mod tests {
}
#[tokio::test]
async fn undeclared_names_still_read_the_process_environment() {
async fn names_excluded_by_hosted_keys_read_the_process_environment() {
let name = "LITELLM_RUST_BRIDGE_UNDECLARED";
unsafe { std::env::set_var(name, "env-key") };
let server = MockServer::start().await;

View file

@ -0,0 +1,463 @@
use std::{collections::BTreeMap, sync::Arc};
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
use litellm_host_python::{from_py, json_object_field, run_async_value, run_sync_value, to_py};
use litellm_secrets::{
KeyManagementSettings, KeyManagementSystem, Secret, SecretManager, load_native_manager,
read_secret_from_python_manager,
};
use litellm_secrets_types::PythonSecretRead;
use pyo3::{
exceptions::{PyAttributeError, PyRuntimeError, PyValueError},
prelude::*,
};
#[derive(Clone, PartialEq)]
struct Configuration {
system: KeyManagementSystem,
settings: KeyManagementSettings,
environment: BTreeMap<String, String>,
enterprise_enabled: bool,
}
#[pyclass(frozen, name = "_SecretManagerRuntime")]
pub(crate) struct NativeSecretManager {
backend: SecretManager,
configuration: Configuration,
pid: u32,
}
impl NativeSecretManager {
pub(super) fn backend(&self) -> PyResult<SecretManager> {
if self.pid != std::process::id() {
return Err(PyRuntimeError::new_err(
"native secret manager must be recreated after fork",
));
}
Ok(self.backend.clone())
}
fn build(py: Python<'_>, configuration: Configuration) -> PyResult<Self> {
let values = configuration.environment.clone();
let environment: Arc<dyn Lookup + Send + Sync> =
Arc::new(move |name: &str| values.get(name).cloned());
let system = configuration.system;
let settings = configuration.settings.clone();
let enterprise_enabled = configuration.enterprise_enabled;
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()))
})?;
Ok(Self {
backend,
configuration,
pid: std::process::id(),
})
}
}
#[pymethods]
impl NativeSecretManager {
#[staticmethod]
#[pyo3(signature = (system, environment, settings=None, enterprise_enabled=false))]
fn from_config(
py: Python<'_>,
system: &str,
environment: BTreeMap<String, String>,
settings: Option<&Bound<'_, PyAny>>,
enterprise_enabled: bool,
) -> PyResult<Self> {
let system = serde_json::from_value(serde_json::Value::String(system.to_owned()))
.map_err(|_| PyValueError::new_err("unknown secret manager system"))?;
let settings = parse_settings(settings)?;
Self::build(
py,
Configuration {
system,
settings,
environment,
enterprise_enabled,
},
)
}
#[staticmethod]
pub(super) fn from_client(client: &Bound<'_, PyAny>) -> PyResult<Option<Py<Self>>> {
let py = client.py();
if let Ok(native) = client.extract::<Py<Self>>() {
native.borrow(py).backend()?;
return Ok(Some(native));
}
let config = py
.import("litellm.rust_bridge.secret_manager")?
.getattr("native_secret_manager_config")?
.call1((client,))?;
if config.is_none() {
return Ok(None);
}
if !config.getattr("owner_type")?.is(client.get_type()) {
return Ok(None);
}
let methods = config
.getattr("methods")?
.extract::<Vec<(String, Py<PyAny>)>>()?;
for (name, original) in methods {
let current = client.getattr(name.as_str())?;
let implementation = optional_attribute(&current, "__func__")?.unwrap_or(current);
if !implementation.is(original.bind(py)) {
return Ok(None);
}
}
let environment_attributes: BTreeMap<String, String> = config
.getattr("environment_attributes")?
.extract::<Vec<(String, String)>>()?
.into_iter()
.collect();
let captured = config
.getattr("environment")?
.extract::<Vec<(String, String)>>()?;
let overrides = environment_attributes
.iter()
.map(|(key, attribute)| {
let value = attribute_path(client, attribute)?;
Ok((
key.clone(),
if value.is_none() {
None
} else {
Some(value.str()?.extract::<String>()?)
},
))
})
.collect::<PyResult<Vec<_>>>()?;
let settings =
from_py::<serde_json::Map<String, serde_json::Value>>(&config.getattr("settings")?)
.map_err(|_| PyValueError::new_err("invalid secret manager settings"))?;
let attributes = config
.getattr("settings_attributes")?
.extract::<Vec<String>>()?;
let setting_overrides = attributes
.into_iter()
.map(|name| {
let value = from_py::<serde_json::Value>(&client.getattr(name.as_str())?)?;
Ok((name, value))
})
.collect::<PyResult<Vec<_>>>()?;
let configuration = Configuration {
system: serde_json::from_value(serde_json::Value::String(
config.getattr("system")?.extract()?,
))
.map_err(|_| PyValueError::new_err("unknown secret manager system"))?,
settings: serde_json::from_value(serde_json::Value::Object(
settings.into_iter().chain(setting_overrides).collect(),
))
.map_err(|_| PyValueError::new_err("invalid secret manager settings"))?,
environment: captured
.into_iter()
.filter(|(key, _)| !environment_attributes.contains_key(key))
.chain(
overrides
.into_iter()
.filter_map(|(key, value)| value.map(|value| (key, value))),
)
.collect(),
enterprise_enabled: config.getattr("enterprise_enabled")?.extract()?,
};
if let Some(native) = cached(client, &configuration)? {
return Ok(Some(native));
}
let runtime = Self::build(py, configuration)?;
if let Some(native) = cached(client, &runtime.configuration)? {
return Ok(Some(native));
}
let native = Py::new(py, runtime)?;
client.setattr("_litellm_native_secret_manager", native.bind(py))?;
Ok(Some(native))
}
#[getter]
fn system(&self) -> String {
serde_json::to_value(self.configuration.system)
.expect("serializable system")
.as_str()
.expect("string system")
.to_owned()
}
#[pyo3(signature = (name, settings=None))]
fn read_secret(
&self,
py: Python<'_>,
name: String,
settings: Option<&Bound<'_, PyAny>>,
) -> PyResult<Py<PyAny>> {
let backend = self.backend()?;
let settings = settings
.map(|value| parse_settings(Some(value)))
.transpose()?
.unwrap_or_else(|| self.configuration.settings.clone());
run_sync_value(py, async move {
read_secret_from_python_manager(&backend, &name, &settings, &ProcessEnvironment)
.await
.map_err(|error| python_read_error(backend.system(), &name, error))
.and_then(|value| python_secret_value(value, &name))
})
}
#[pyo3(signature = (secret_name, optional_params=None, timeout=None, primary_secret_name=None))]
fn sync_read_secret(
&self,
py: Python<'_>,
secret_name: String,
optional_params: Option<&Bound<'_, PyAny>>,
timeout: Option<&Bound<'_, PyAny>>,
primary_secret_name: Option<String>,
) -> PyResult<Py<PyAny>> {
let backend = self.backend()?;
let request = super::provider::read_request(
self.configuration.system,
secret_name,
optional_params,
timeout,
primary_secret_name,
true,
)?;
run_sync_value(py, async move {
super::operations::read_python_provider(&backend, &request, &ProcessEnvironment)
.await
.map_err(|error| PyValueError::new_err(error.to_string()))
.and_then(|value| python_secret_value(value, &request.secret_name))
})
}
#[pyo3(signature = (secret_name, optional_params=None, timeout=None, primary_secret_name=None))]
fn async_read_secret<'py>(
&self,
py: Python<'py>,
secret_name: String,
optional_params: Option<&Bound<'py, PyAny>>,
timeout: Option<&Bound<'py, PyAny>>,
primary_secret_name: Option<String>,
) -> PyResult<Bound<'py, PyAny>> {
let backend = self.backend()?;
let request = super::provider::read_request(
self.configuration.system,
secret_name,
optional_params,
timeout,
primary_secret_name,
false,
)?;
run_async_value(py, async move {
super::operations::read_python_provider(&backend, &request, &ProcessEnvironment)
.await
.map_err(|error| PyValueError::new_err(error.to_string()))
.and_then(|value| python_secret_value(value, &request.secret_name))
})
}
#[pyo3(signature = (secret_name, secret_value, description=None, optional_params=None, timeout=None, tags=None))]
#[expect(
clippy::too_many_arguments,
reason = "preserves the Python secret-manager write signature"
)]
fn async_write_secret<'py>(
&self,
py: Python<'py>,
secret_name: String,
secret_value: String,
description: Option<&Bound<'py, PyAny>>,
optional_params: Option<&Bound<'py, PyAny>>,
timeout: Option<&Bound<'py, PyAny>>,
tags: Option<&Bound<'py, PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
let backend = self.backend()?;
let _ = tags;
let context = litellm_secrets_types::SecretWriteContext {
operation: super::provider::mutation_context(
self.configuration.system,
optional_params,
timeout,
)?,
description: if self.configuration.system == KeyManagementSystem::HashicorpVault {
description
.filter(|value| !value.is_none())
.map(|value| {
if value.is_truthy()? {
value.extract().map(Some)
} else {
Ok(None)
}
})
.transpose()?
.flatten()
} else {
None
},
..litellm_secrets_types::SecretWriteContext::default()
};
let error_context =
super::vault::ErrorContext::capture(py, self.configuration.system, timeout)?;
run_async_value(py, async move {
super::mutation::mutation_value(
super::operations::write_python_provider_with_context(
&backend,
&secret_name,
&litellm_secrets::SecretValue::new(secret_value),
&context,
)
.await,
&error_context,
)
})
}
#[pyo3(signature = (secret_name, recovery_window_in_days=None, optional_params=None, timeout=None))]
fn async_delete_secret<'py>(
&self,
py: Python<'py>,
secret_name: String,
recovery_window_in_days: Option<&Bound<'py, PyAny>>,
optional_params: Option<&Bound<'py, PyAny>>,
timeout: Option<&Bound<'py, PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
let backend = self.backend()?;
let _ = recovery_window_in_days;
let context =
super::provider::mutation_context(self.configuration.system, optional_params, timeout)?;
let error_context =
super::vault::ErrorContext::capture(py, self.configuration.system, timeout)?;
run_async_value(py, async move {
super::mutation::mutation_value(
super::operations::delete_python_provider_with_context(
&backend,
&secret_name,
&context,
)
.await,
&error_context,
)
})
}
#[pyo3(signature = (current_secret_name, new_secret_name, new_secret_value, optional_params=None, timeout=None))]
fn async_rotate_secret<'py>(
&self,
py: Python<'py>,
current_secret_name: String,
new_secret_name: String,
new_secret_value: String,
optional_params: Option<&Bound<'py, PyAny>>,
timeout: Option<&Bound<'py, PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
let backend = self.backend()?;
let context =
super::provider::mutation_context(self.configuration.system, optional_params, timeout)?;
let error_context =
super::vault::ErrorContext::capture(py, self.configuration.system, timeout)?;
run_async_value(py, async move {
super::mutation::mutation_value(
super::operations::rotate_python_provider_with_context(
&backend,
&current_secret_name,
&new_secret_name,
&litellm_secrets::SecretValue::new(new_secret_value),
&context,
)
.await,
&error_context,
)
})
}
#[pyo3(signature = (name, settings=None))]
fn read_secret_async<'py>(
&self,
py: Python<'py>,
name: String,
settings: Option<&Bound<'py, PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
let backend = self.backend()?;
let settings = settings
.map(|value| parse_settings(Some(value)))
.transpose()?
.unwrap_or_else(|| self.configuration.settings.clone());
run_async_value(py, async move {
read_secret_from_python_manager(&backend, &name, &settings, &ProcessEnvironment)
.await
.map_err(|error| python_read_error(backend.system(), &name, error))
.and_then(|value| python_secret_value(value, &name))
})
}
}
fn optional_attribute<'py>(
object: &Bound<'py, PyAny>,
name: &str,
) -> PyResult<Option<Bound<'py, PyAny>>> {
match object.getattr(name) {
Ok(value) => Ok(Some(value)),
Err(error) if error.is_instance_of::<PyAttributeError>(object.py()) => Ok(None),
Err(error) => Err(error),
}
}
fn attribute_path<'py>(object: &Bound<'py, PyAny>, path: &str) -> PyResult<Bound<'py, PyAny>> {
match path.split_once('.') {
Some((head, tail)) => attribute_path(&object.getattr(head)?, tail),
None => object.getattr(path),
}
}
fn cached(
client: &Bound<'_, PyAny>,
configuration: &Configuration,
) -> PyResult<Option<Py<NativeSecretManager>>> {
let Some(value) = optional_attribute(client, "_litellm_native_secret_manager")? else {
return Ok(None);
};
let native = value.extract::<Py<NativeSecretManager>>()?;
let same_configuration = native.borrow(client.py()).pid == std::process::id()
&& &native.borrow(client.py()).configuration == configuration;
Ok(same_configuration.then_some(native))
}
fn parse_settings(value: Option<&Bound<'_, PyAny>>) -> PyResult<KeyManagementSettings> {
value
.map(|value| {
serde_json::from_value(from_py::<serde_json::Value>(value)?)
.map_err(|_| PyValueError::new_err("invalid secret manager settings"))
})
.transpose()
.map(Option::unwrap_or_default)
}
fn python_secret_value(payload: PythonSecretRead, name: &str) -> PyResult<Py<PyAny>> {
let value = match payload {
PythonSecretRead::Value(value) => value,
PythonSecretRead::PrimaryJson(document) => {
return Python::attach(|py| json_object_field(py, document.expose(), name));
}
};
let value = match value {
None => serde_json::Value::Null,
Some(Secret::String(value)) => serde_json::Value::String(value.expose().to_owned()),
Some(Secret::Bool(value)) => serde_json::Value::Bool(value),
Some(Secret::Json(value)) => value,
};
Python::attach(|py| to_py(py, &value))
}
fn python_read_error(
system: KeyManagementSystem,
name: &str,
error: litellm_secrets::Error,
) -> PyErr {
let message = match (system, error) {
(KeyManagementSystem::Cyberark, litellm_secrets::Error::ManagedSecretMissing) => {
format!("No secret found in CyberArk Secret Manager for {name}")
}
(_, error) => error.to_string(),
};
PyValueError::new_err(message)
}

View file

@ -0,0 +1,182 @@
mod operation;
pub(super) use operation::{Failure, FailureKind, FailureStage, delete, rotate, write};
use litellm_secrets::hashicorp::{Error, RawOperationError};
use pyo3::prelude::*;
use super::mutation::{error_value, http_message, json_value};
pub(super) fn failure_value(
py: Python<'_>,
failure: Failure,
context: &ErrorContext,
) -> PyResult<Py<PyAny>> {
let message = match *failure.kind {
FailureKind::Native(RawOperationError::Http {
method,
url,
status,
body,
}) => match failure.stage {
FailureStage::Current(name) if status == 404 => {
format!("Current secret {name} not found")
}
FailureStage::Replacement(name) if status == 404 => {
format!("Failed to verify new secret {name}")
}
FailureStage::Current(_) => format!(
"HTTP error occurred while checking current secret: {}",
response_text(py, &body)?
),
FailureStage::Replacement(_) => format!(
"HTTP error occurred while verifying new secret: {}",
response_text(py, &body)?
),
FailureStage::Mutation => http_message(py, &method, &url, status)?,
},
FailureKind::ValueMismatch { expected, actual } => {
let actual = json_value(py, &actual)?;
format!(
"New secret value mismatch. Expected: {}, Got: {}",
expected.expose(),
actual.bind(py).str()?
)
}
kind => {
let message = cause_message(py, kind, context)?;
match failure.stage {
FailureStage::Current(_) => {
format!("Error checking current secret: {message}")
}
FailureStage::Replacement(_) => {
format!("Error verifying new secret: {message}")
}
FailureStage::Mutation => message,
}
}
};
error_value(py, message)
}
fn cause_message(py: Python<'_>, kind: FailureKind, context: &ErrorContext) -> PyResult<String> {
Ok(match kind {
FailureKind::Native(RawOperationError::Local(error)) => error.to_string(),
FailureKind::UnsafeName(name) => {
format!("Invalid secret_name {}", name.into_pyobject(py)?.repr()?)
}
FailureKind::Native(RawOperationError::Timeout { method, elapsed }) => {
if method == "POST" {
let elapsed = py
.import("builtins")?
.call_method1("round", (elapsed.as_secs_f64(), 3))?;
let kwargs = pyo3::types::PyDict::new(py);
kwargs.set_item(
"message",
format!(
"Connection timed out. Timeout passed={}, time taken={} seconds",
context.timeout.as_deref().unwrap_or("None"),
elapsed.str()?
),
)?;
kwargs.set_item("model", "default-model-name")?;
kwargs.set_item("llm_provider", "litellm-httpx-handler")?;
kwargs.set_item("headers", pyo3::types::PyDict::new(py))?;
py.import("litellm")?
.getattr("Timeout")?
.call((), Some(&kwargs))?
.str()?
.extract()?
} else if context.aiohttp {
"Timeout on reading data from socket".to_owned()
} else {
String::new()
}
}
FailureKind::Native(RawOperationError::Transport(source)) => {
if let Some(error) = request_error(&source) {
if error.is_timeout() {
String::new()
} else if error.is_connect() {
"All connection attempts failed".to_owned()
} else {
"HashiCorp Vault request failed".to_owned()
}
} else {
"HashiCorp Vault request failed".to_owned()
}
}
FailureKind::MissingGet(value) => {
let value = json_value(py, &value)?;
match value.bind(py).getattr("get") {
Err(error) => error.value(py).str()?.extract()?,
Ok(_) => "HashiCorp Vault response payload is malformed".to_owned(),
}
}
FailureKind::Json(body) => match json_value(py, &body) {
Err(error) => error.value(py).str()?.extract()?,
Ok(_) => "HashiCorp Vault response payload is malformed".to_owned(),
},
FailureKind::Native(RawOperationError::Authentication {
source,
url,
certificate,
}) => {
let message = match source {
Error::LoginStatus { status } => http_message(py, "POST", &url, status)?,
error => error.to_string(),
};
let mechanism = if certificate { "TLS cert" } else { "AppRole" };
format!("Could not authenticate to Vault via {mechanism}: {message}")
}
FailureKind::Native(RawOperationError::Http {
method,
url,
status,
..
}) => http_message(py, &method, &url, status)?,
FailureKind::ValueMismatch { .. } => "New secret value mismatch".to_owned(),
})
}
fn request_error<'a>(error: &'a (dyn std::error::Error + 'static)) -> Option<&'a reqwest::Error> {
error
.downcast_ref::<reqwest::Error>()
.or_else(|| error.source().and_then(request_error))
}
fn response_text(py: Python<'_>, body: &[u8]) -> PyResult<String> {
let kwargs = pyo3::types::PyDict::new(py);
kwargs.set_item("content", pyo3::types::PyBytes::new(py, body))?;
py.import("httpx")?
.getattr("Response")?
.call((200,), Some(&kwargs))?
.getattr("text")?
.extract()
}
#[derive(Default)]
pub(super) struct ErrorContext {
timeout: Option<String>,
aiohttp: bool,
}
impl ErrorContext {
pub(super) fn capture(
py: Python<'_>,
system: litellm_secrets::KeyManagementSystem,
timeout: Option<&Bound<'_, PyAny>>,
) -> PyResult<Self> {
if system != litellm_secrets::KeyManagementSystem::HashicorpVault {
return Ok(Self::default());
}
Ok(Self {
timeout: timeout.map(|value| value.str()?.extract()).transpose()?,
aiohttp: py
.import("litellm.llms.custom_httpx.http_handler")?
.getattr("AsyncHTTPHandler")?
.call_method0("_should_use_aiohttp_transport")?
.extract()?,
})
}
}

View file

@ -0,0 +1,168 @@
use std::collections::HashMap;
use litellm_secrets::{
SecretValue,
hashicorp::{Error, HashicorpVault, RawOperationError},
};
use litellm_secrets_types::{HashicorpOperationContext, SecretWriteContext};
use serde_json::value::RawValue;
#[derive(Debug)]
pub(crate) enum FailureStage {
Mutation,
Current(String),
Replacement(String),
}
#[derive(veil::Redact)]
pub(crate) enum FailureKind {
Native(RawOperationError),
UnsafeName(#[redact] String),
Json(#[redact] Vec<u8>),
MissingGet(#[redact] Vec<u8>),
ValueMismatch {
expected: SecretValue,
#[redact]
actual: Vec<u8>,
},
}
#[derive(Debug)]
pub(crate) struct Failure {
pub kind: Box<FailureKind>,
pub stage: FailureStage,
}
impl From<FailureKind> for Failure {
fn from(kind: FailureKind) -> Self {
Self {
kind: Box::new(kind),
stage: FailureStage::Mutation,
}
}
}
impl Failure {
fn during(self, stage: FailureStage) -> Self {
if matches!(*self.kind, FailureKind::UnsafeName(_)) {
self
} else {
Self { stage, ..self }
}
}
}
fn native_failure(name: &str, error: RawOperationError) -> Failure {
match error {
RawOperationError::Local(Error::InvalidSecretName(_)) => {
FailureKind::UnsafeName(name.to_owned()).into()
}
error => FailureKind::Native(error).into(),
}
}
pub(crate) async fn write(
client: &HashicorpVault,
name: &str,
value: &SecretValue,
context: &SecretWriteContext<HashicorpOperationContext>,
) -> Result<Vec<u8>, Failure> {
client
.write_raw(name, value, context)
.await
.map_err(|error| native_failure(name, error))
}
pub(crate) async fn delete(
client: &HashicorpVault,
name: &str,
context: &HashicorpOperationContext,
) -> Result<(), Failure> {
client
.delete_raw(name, context)
.await
.map_err(|error| native_failure(name, error))
}
pub(crate) async fn rotate(
client: &HashicorpVault,
current_name: &str,
new_name: &str,
value: &SecretValue,
context: &HashicorpOperationContext,
) -> Result<Vec<u8>, Failure> {
client
.read_raw(current_name, context)
.await
.map_err(|error| {
native_failure(current_name, error)
.during(FailureStage::Current(current_name.to_owned()))
})?;
let response = write(
client,
new_name,
value,
&SecretWriteContext {
description: Some(format!("Rotated from {current_name}")),
operation: context.clone(),
..SecretWriteContext::default()
},
)
.await?;
let parsed: &RawValue =
serde_json::from_slice(&response).map_err(|_| FailureKind::Json(response.clone()))?;
let status = raw_object_field(parsed, "status")
.ok()
.flatten()
.and_then(|value| serde_json::from_slice::<String>(value.get().as_bytes()).ok());
if status.as_deref() == Some("error") {
return Ok(response);
}
let verification = client.read_raw(new_name, context).await.map_err(|error| {
native_failure(new_name, error).during(FailureStage::Replacement(new_name.to_owned()))
})?;
let parsed: &RawValue = serde_json::from_slice(&verification).map_err(|_| {
Failure::from(FailureKind::Json(verification.clone()))
.during(FailureStage::Replacement(new_name.to_owned()))
})?;
let data_key = context
.data_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.unwrap_or("key");
let actual = verification_value(parsed, data_key).map_err(|failure| {
Failure::from(failure).during(FailureStage::Replacement(new_name.to_owned()))
})?;
let actual_string = serde_json::from_slice::<String>(actual.get().as_bytes()).ok();
if actual_string.as_deref() != Some(value.expose()) {
return Err(FailureKind::ValueMismatch {
expected: value.clone(),
actual: actual.get().as_bytes().to_vec(),
}
.into());
}
if current_name != new_name {
let _ = delete(client, current_name, context).await;
}
Ok(response)
}
fn verification_value<'a>(document: &'a RawValue, key: &str) -> Result<&'a RawValue, FailureKind> {
let Some(outer) = raw_object_field(document, "data")? else {
return Ok(RawValue::NULL);
};
let Some(inner) = raw_object_field(outer, "data")? else {
return Ok(RawValue::NULL);
};
Ok(raw_object_field(inner, key)?.unwrap_or(RawValue::NULL))
}
fn raw_object_field<'a>(
document: &'a RawValue,
key: &str,
) -> Result<Option<&'a RawValue>, FailureKind> {
let object: HashMap<String, &RawValue> = serde_json::from_slice(document.get().as_bytes())
.map_err(|_| FailureKind::MissingGet(document.get().as_bytes().to_vec()))?;
Ok(object.get(key).copied())
}

View file

@ -0,0 +1 @@
- https://docs.aws.amazon.com/secretsmanager/latest/apireference/Welcome.html

View file

@ -22,3 +22,4 @@ base64.workspace = true
rstest.workspace = true
tokio.workspace = true
wiremock = "0.6.5"
tempfile = "3"

View file

@ -7,7 +7,7 @@ use litellm_auth_aws::{
resolve_credentials,
};
use litellm_core_utils::settings::Lookup;
use litellm_secrets_types::KeyManagementSettings;
use litellm_secrets_types::{AwsOperationContext, KeyManagementSettings};
use crate::Error;
@ -21,9 +21,29 @@ impl Credentials {
pub(crate) fn new(
settings: &KeyManagementSettings,
environment: Arc<dyn Lookup + Send + Sync>,
) -> Self {
Self::with_context(settings, environment, &AwsOperationContext::default())
}
pub(crate) fn with_context(
settings: &KeyManagementSettings,
environment: Arc<dyn Lookup + Send + Sync>,
context: &AwsOperationContext,
) -> Self {
Self {
config: AwsAuthConfig {
access_key_id: context
.access_key_id
.as_ref()
.map(|value| value.expose().to_owned()),
secret_access_key: context
.secret_access_key
.as_ref()
.map(|value| value.expose().to_owned()),
session_token: context
.session_token
.as_ref()
.map(|value| value.expose().to_owned()),
region_name: region(settings, environment.as_ref()).ok(),
role_name: settings.aws_role_name.clone(),
session_name: settings.aws_session_name.clone(),
@ -37,7 +57,6 @@ impl Credentials {
.as_ref()
.map(|v| v.expose().to_owned()),
sts_endpoint: settings.aws_sts_endpoint.clone(),
..Default::default()
},
environment,
}

View file

@ -6,8 +6,6 @@ pub enum Error {
Auth(#[from] #[redact] litellm_auth_aws::Error),
#[error("AWS region is not configured")]
MissingRegion,
#[error("AWS Secrets Manager received a non-AWS operation context")]
InvalidOperationContext,
#[error("AWS Secrets Manager was constructed without context-aware configuration")]
OperationContextUnavailable,
#[error("KMS response has no plaintext")]
@ -20,6 +18,12 @@ pub enum Error {
Read(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::get_secret_value::GetSecretValueError>>),
#[error("AWS Secrets Manager create failed")]
Create(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::create_secret::CreateSecretError>>),
#[error("AWS Secrets Manager restore failed")]
Restore(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::restore_secret::RestoreSecretError>>),
#[error("AWS Secrets Manager restored update failed")]
Update(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::update_secret::UpdateSecretError>>),
#[error("AWS Secrets Manager tagging failed")]
Tag(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::tag_resource::TagResourceError>>),
#[error("AWS Secrets Manager update failed")]
Put(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::put_secret_value::PutSecretValueError>>),
#[error("AWS Secrets Manager delete failed")]

View file

@ -1,4 +1,3 @@
use litellm_auth_aws::constants::AWS_REGION_NAME;
use std::sync::Arc;
use aws_sdk_kms::{
@ -37,10 +36,7 @@ impl AwsKms {
}
pub fn validate_environment(environment: &dyn Lookup) -> Result<(), Error> {
environment
.get(AWS_REGION_NAME)
.map(|_| ())
.ok_or(Error::MissingRegion)
auth::region(&KeyManagementSettings::default(), environment).map(|_| ())
}
pub fn load_aws_kms(
@ -51,9 +47,6 @@ pub fn load_aws_kms(
if use_aws_kms != Some(true) {
return Ok(None);
}
if settings.aws_region_name.is_none() {
validate_environment(environment.as_ref())?;
}
let config = aws_sdk_kms::Config::builder()
.behavior_version(BehaviorVersion::latest())
.region(Region::new(auth::region(settings, environment.as_ref())?))

View file

@ -1,3 +1,9 @@
mod client;
mod read;
mod write;
pub use read::is_bootstrap_key;
use litellm_auth_aws::constants::AWS_BEDROCK_RUNTIME_ENDPOINT;
use std::{collections::BTreeMap, sync::Arc};
@ -16,8 +22,9 @@ use litellm_auth_aws::constants::{
};
use litellm_core_utils::settings::Lookup;
use litellm_secrets_types::{
AwsOperationContext, BaseSecretManager, KeyManagementSettings, Secret, SecretOperationContext,
SecretValue, SecretWriteContext, async_rotate_secret,
AwsOperationContext, BaseSecretManager, KeyManagementSettings, RotationError, Secret,
SecretDeleter, SecretRotator, SecretValue, SecretWriteContext, SecretWriter,
async_rotate_secret,
};
use serde_json::Value;
@ -68,395 +75,4 @@ impl AwsSecretsManagerV2 {
write_settings,
}
}
fn with_context_client_factory(
client: Client,
write_settings: AwsSecretWriteSettings,
context_client_factory: ContextClientFactory,
) -> Self {
Self {
client,
context_client_factory: Some(Box::new(context_client_factory)),
write_settings,
}
}
pub fn load_aws_secret_manager(
use_aws_secret_manager: Option<bool>,
settings: KeyManagementSettings,
environment: Arc<dyn Lookup + Send + Sync>,
) -> Result<Option<Self>, Error> {
if use_aws_secret_manager != Some(true) {
return Ok(None);
}
let context_client_factory = ContextClientFactory {
settings: settings.clone(),
environment: environment.clone(),
endpoint_url: environment
.get(AWS_BEDROCK_RUNTIME_ENDPOINT)
.map(|url| url.replace("bedrock-runtime", "secretsmanager")),
};
let client = context_client_factory.client(&AwsOperationContext::default())?;
Ok(Some(Self::with_context_client_factory(
client,
(&settings).into(),
context_client_factory,
)))
}
pub async fn read_secret_for_resolver(
&self,
name: &str,
primary_name: Option<&str>,
environment: &(dyn Lookup + Sync),
) -> Result<Option<Secret>, Error> {
if bootstrap_key(name) {
return Ok(environment
.get(name)
.map(SecretValue::new)
.map(Secret::String));
}
match primary_name.filter(|name| !name.is_empty()) {
None => self
.async_read_secret(name)
.await
.map(|value| value.map(Secret::String)),
Some(primary) => {
let value = if bootstrap_key(primary) {
environment.get(primary).map(SecretValue::new)
} else {
self.async_read_secret(primary).await?
};
let Some(value) = value else {
return Ok(None);
};
let object: Value =
serde_json::from_str(value.expose()).map_err(|_| Error::PrimarySecret)?;
let object = object.as_object().ok_or(Error::PrimarySecret)?;
Ok(object.get(name).cloned().map(Secret::from_json))
}
}
}
pub async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
Self::async_read_secret_with_client(&self.client, name).await
}
async fn async_read_secret_with_client(
client: &Client,
name: &str,
) -> Result<Option<SecretValue>, Error> {
match client.get_secret_value().secret_id(name).send().await {
Ok(response) => response
.secret_string
.map(SecretValue::new)
.map(Some)
.ok_or(Error::MissingString),
Err(error)
if matches!(
&error,
aws_sdk_secretsmanager::error::SdkError::TimeoutError(_)
) || matches!(&error, aws_sdk_secretsmanager::error::SdkError::DispatchFailure(failure) if failure.is_timeout()) =>
{
Err(Error::Timeout)
}
Err(error)
if error
.as_service_error()
.is_some_and(|error| error.is_resource_not_found_exception()) =>
{
Ok(None)
}
Err(error) => Err(Error::Read(Box::new(error))),
}
}
pub async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
description: Option<&str>,
) -> Result<CreateSecretOutput, Error> {
self.async_write_secret_with_client_and_tags(&self.client, name, value, description, None)
.await
}
async fn async_write_secret_with_client_and_tags(
&self,
client: &Client,
name: &str,
value: &SecretValue,
description: Option<&str>,
tags: Option<&BTreeMap<String, String>>,
) -> Result<CreateSecretOutput, Error> {
let response = client
.create_secret()
.name(name)
.secret_string(value.expose())
.set_description(description.filter(|v| !v.is_empty()).map(str::to_owned))
.set_kms_key_id(
self.write_settings
.kms_key_id
.clone()
.filter(|v| !v.is_empty()),
)
.set_tags(tags.or(self.write_settings.tags.as_ref()).map(|tags| {
tags.iter()
.map(|(key, value)| Tag::builder().key(key).value(value).build())
.collect()
}))
.send()
.await
.map_err(|error| Error::Create(Box::new(error)))?;
if let Some(regions) = &self.write_settings.replica_regions
&& !regions.is_empty()
&& self
.async_replicate_secret_with_client(client, name, regions)
.await
.is_err()
{
litellm_tracing::warn!("secret created but replication failed");
}
Ok(response)
}
pub async fn async_replicate_secret(
&self,
name: &str,
regions: &[String],
) -> Result<Option<ReplicateSecretToRegionsOutput>, Error> {
self.async_replicate_secret_with_client(&self.client, name, regions)
.await
}
async fn async_replicate_secret_with_client(
&self,
client: &Client,
name: &str,
regions: &[String],
) -> Result<Option<ReplicateSecretToRegionsOutput>, Error> {
if regions.is_empty() {
return Ok(None);
}
client
.replicate_secret_to_regions()
.secret_id(name)
.set_add_replica_regions(Some(
regions
.iter()
.map(|region| ReplicaRegionType::builder().region(region).build())
.collect(),
))
.send()
.await
.map(Some)
.map_err(|error| Error::Replicate(Box::new(error)))
}
pub async fn async_put_secret_value(
&self,
name: &str,
value: &SecretValue,
) -> Result<PutSecretValueOutput, Error> {
self.async_put_secret_value_with_client(&self.client, name, value)
.await
}
async fn async_put_secret_value_with_client(
&self,
client: &Client,
name: &str,
value: &SecretValue,
) -> Result<PutSecretValueOutput, Error> {
client
.put_secret_value()
.secret_id(name)
.secret_string(value.expose())
.send()
.await
.map_err(|error| Error::Put(Box::new(error)))
}
pub async fn async_delete_secret(
&self,
name: &str,
recovery_window_in_days: Option<u32>,
) -> Result<DeleteSecretOutput, Error> {
self.async_delete_secret_with_client(&self.client, name, recovery_window_in_days)
.await
}
async fn async_delete_secret_with_client(
&self,
client: &Client,
name: &str,
recovery_window_in_days: Option<u32>,
) -> Result<DeleteSecretOutput, Error> {
client
.delete_secret()
.secret_id(name)
.set_recovery_window_in_days(recovery_window_in_days.map(i64::from))
.send()
.await
.map_err(|error| Error::Delete(Box::new(error)))
}
pub async fn async_rotate_secret(
&self,
current_name: &str,
new_name: &str,
value: &SecretValue,
) -> Result<RotationResponse, Error> {
self.async_rotate_secret_with_context(
current_name,
new_name,
value,
&SecretOperationContext::default(),
)
.await
}
pub async fn async_rotate_secret_with_context(
&self,
current_name: &str,
new_name: &str,
value: &SecretValue,
context: &SecretOperationContext,
) -> Result<RotationResponse, Error> {
if current_name == new_name {
let client = self.client_for_context(context)?;
return self
.async_put_secret_value_with_client(&client, current_name, value)
.await
.map(RotationResponse::Updated);
}
async_rotate_secret(self, current_name, new_name, value, context)
.await
.map(RotationResponse::Created)
}
fn client_for_context(&self, context: &SecretOperationContext) -> Result<Client, Error> {
match context {
SecretOperationContext::Default => Ok(self.client.clone()),
SecretOperationContext::Aws(context) if context == &AwsOperationContext::default() => {
Ok(self.client.clone())
}
SecretOperationContext::Aws(context) => self
.context_client_factory
.as_ref()
.ok_or(Error::OperationContextUnavailable)?
.client(context),
_ => Err(Error::InvalidOperationContext),
}
}
}
impl ContextClientFactory {
fn client(&self, context: &AwsOperationContext) -> Result<Client, Error> {
let settings = KeyManagementSettings {
aws_region_name: context
.region_name
.clone()
.or_else(|| self.settings.aws_region_name.clone()),
aws_role_name: context
.role_name
.clone()
.or_else(|| self.settings.aws_role_name.clone()),
aws_session_name: context
.session_name
.clone()
.or_else(|| self.settings.aws_session_name.clone()),
aws_external_id: context
.external_id
.clone()
.or_else(|| self.settings.aws_external_id.clone()),
aws_profile_name: context
.profile_name
.clone()
.or_else(|| self.settings.aws_profile_name.clone()),
aws_web_identity_token: context
.web_identity_token
.clone()
.or_else(|| self.settings.aws_web_identity_token.clone()),
aws_sts_endpoint: context
.sts_endpoint
.clone()
.or_else(|| self.settings.aws_sts_endpoint.clone()),
..self.settings.clone()
};
let builder = aws_sdk_secretsmanager::Config::builder()
.behavior_version(BehaviorVersion::latest())
.region(Region::new(auth::region(
&settings,
self.environment.as_ref(),
)?))
.credentials_provider(auth::Credentials::new(&settings, self.environment.clone()));
let builder = match context.timeout {
Some(timeout) => builder.timeout_config(
aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder()
.operation_timeout(timeout)
.build(),
),
None => builder,
};
let config = match &self.endpoint_url {
Some(endpoint_url) => builder.endpoint_url(endpoint_url.clone()).build(),
None => builder.build(),
};
Ok(Client::from_conf(config))
}
}
impl BaseSecretManager for AwsSecretsManagerV2 {
type Error = Error;
type WriteResponse = CreateSecretOutput;
type DeleteResponse = DeleteSecretOutput;
async fn async_read_secret(
&self,
name: &str,
context: &SecretOperationContext,
) -> Result<Option<SecretValue>, Error> {
let client = self.client_for_context(context)?;
Self::async_read_secret_with_client(&client, name).await
}
async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
context: &SecretWriteContext,
) -> Result<CreateSecretOutput, Error> {
let client = self.client_for_context(&context.operation)?;
self.async_write_secret_with_client_and_tags(
&client,
name,
value,
context.description.as_deref(),
(!context.tags.is_empty()).then_some(&context.tags),
)
.await
}
async fn async_delete_secret(
&self,
name: &str,
recovery_window_in_days: Option<u32>,
context: &SecretOperationContext,
) -> Result<DeleteSecretOutput, Error> {
let client = self.client_for_context(context)?;
self.async_delete_secret_with_client(&client, name, recovery_window_in_days)
.await
}
}
fn bootstrap_key(name: &str) -> bool {
matches!(
name,
AWS_ACCESS_KEY_ID
| AWS_SECRET_ACCESS_KEY
| AWS_REGION_NAME
| AWS_REGION
| AWS_BEDROCK_RUNTIME_ENDPOINT
)
}

View file

@ -0,0 +1,116 @@
use super::*;
impl AwsSecretsManagerV2 {
pub(super) fn with_context_client_factory(
client: Client,
write_settings: AwsSecretWriteSettings,
context_client_factory: ContextClientFactory,
) -> Self {
Self {
client,
context_client_factory: Some(Box::new(context_client_factory)),
write_settings,
}
}
pub fn load_aws_secret_manager(
use_aws_secret_manager: Option<bool>,
settings: KeyManagementSettings,
environment: Arc<dyn Lookup + Send + Sync>,
) -> Result<Option<Self>, Error> {
if use_aws_secret_manager != Some(true) {
return Ok(None);
}
let context_client_factory = ContextClientFactory {
settings: settings.clone(),
environment: environment.clone(),
endpoint_url: environment
.get(AWS_BEDROCK_RUNTIME_ENDPOINT)
.map(|url| url.replace("bedrock-runtime", "secretsmanager")),
};
let client = context_client_factory.client(&AwsOperationContext::default())?;
Ok(Some(Self::with_context_client_factory(
client,
(&settings).into(),
context_client_factory,
)))
}
pub(super) fn client_for_context(
&self,
context: &AwsOperationContext,
) -> Result<Client, Error> {
if context == &AwsOperationContext::default() {
return Ok(self.client.clone());
}
self.context_client_factory
.as_ref()
.ok_or(Error::OperationContextUnavailable)?
.client(context)
}
}
impl ContextClientFactory {
fn client(&self, context: &AwsOperationContext) -> Result<Client, Error> {
let settings = KeyManagementSettings {
aws_region_name: context
.region_name
.clone()
.or_else(|| self.settings.aws_region_name.clone()),
aws_role_name: context
.role_name
.clone()
.or_else(|| self.settings.aws_role_name.clone()),
aws_session_name: context
.session_name
.clone()
.or_else(|| self.settings.aws_session_name.clone()),
aws_external_id: context
.external_id
.clone()
.or_else(|| self.settings.aws_external_id.clone()),
aws_profile_name: context
.profile_name
.clone()
.or_else(|| self.settings.aws_profile_name.clone()),
aws_web_identity_token: context
.web_identity_token
.clone()
.or_else(|| self.settings.aws_web_identity_token.clone()),
aws_sts_endpoint: context
.sts_endpoint
.clone()
.or_else(|| self.settings.aws_sts_endpoint.clone()),
..self.settings.clone()
};
let builder = aws_sdk_secretsmanager::Config::builder()
.behavior_version(BehaviorVersion::latest())
.region(Region::new(auth::region(
&settings,
self.environment.as_ref(),
)?))
.credentials_provider(auth::Credentials::with_context(
&settings,
self.environment.clone(),
context,
));
let builder = match context.timeout {
Some(timeout) => builder.timeout_config(
aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder()
.operation_timeout(timeout)
.build(),
),
None => builder,
};
let endpoint_url = context
.bedrock_runtime_endpoint
.as_ref()
.map(|url| url.replace("bedrock-runtime", "secretsmanager"))
.or_else(|| self.endpoint_url.clone());
let config = match endpoint_url {
Some(endpoint_url) => builder.endpoint_url(endpoint_url).build(),
None => builder.build(),
};
Ok(Client::from_conf(config))
}
}

View file

@ -0,0 +1,225 @@
use super::*;
use aws_sdk_secretsmanager::config::retry::RetryConfig;
use litellm_secrets_types::PythonSecretRead;
#[derive(Clone, Copy)]
enum ReadPolicy {
Native,
Python,
}
impl AwsSecretsManagerV2 {
pub async fn read_secret_for_resolver(
&self,
name: &str,
primary_name: Option<&str>,
environment: &(dyn Lookup + Sync),
) -> Result<Option<Secret>, Error> {
let payload = self
.read_payload(name, primary_name, environment, ReadPolicy::Native)
.await?;
resolve_payload(payload, name)
}
pub async fn read_secret_for_python(
&self,
name: &str,
primary_name: Option<&str>,
environment: &(dyn Lookup + Sync),
) -> Result<Option<Secret>, Error> {
let payload = self
.read_payload_for_python(name, primary_name, environment)
.await?;
resolve_payload(payload, name)
}
pub async fn read_payload_for_python(
&self,
name: &str,
primary_name: Option<&str>,
environment: &(dyn Lookup + Sync),
) -> Result<PythonSecretRead, Error> {
self.read_payload(name, primary_name, environment, ReadPolicy::Python)
.await
}
pub async fn read_provider_payload_for_python(
&self,
name: &str,
primary_name: Option<&str>,
context: &AwsOperationContext,
synchronous: bool,
environment: &(dyn Lookup + Sync),
) -> Result<PythonSecretRead, Error> {
if synchronous && is_bootstrap_key(name) {
return Ok(PythonSecretRead::Value(
environment
.get(name)
.map(SecretValue::new)
.map(Secret::String),
));
}
if let Some(primary) = primary_name.filter(|value| !value.is_empty()) {
let value = if synchronous && is_bootstrap_key(primary) {
environment.get(primary).map(SecretValue::new)
} else {
self.read_with_policy(primary, ReadPolicy::Python).await?
};
return Ok(match value.filter(|value| !value.expose().is_empty()) {
Some(value) => PythonSecretRead::PrimaryJson(value),
None => PythonSecretRead::Value(None),
});
}
let client = self.client_for_context(context)?;
let value = match Self::read_with_client(&client, name, ReadPolicy::Python).await {
Err(Error::Read(_) | Error::MissingString | Error::Timeout) => None,
result => result?,
};
Ok(PythonSecretRead::Value(value.map(Secret::String)))
}
async fn read_payload(
&self,
name: &str,
primary_name: Option<&str>,
environment: &(dyn Lookup + Sync),
policy: ReadPolicy,
) -> Result<PythonSecretRead, Error> {
if is_bootstrap_key(name) {
return Ok(PythonSecretRead::Value(
environment
.get(name)
.map(SecretValue::new)
.map(Secret::String),
));
}
match primary_name.filter(|name| !name.is_empty()) {
None => self
.read_with_policy(name, policy)
.await
.map(|value| PythonSecretRead::Value(value.map(Secret::String))),
Some(primary) => {
let value = if is_bootstrap_key(primary) {
environment.get(primary).map(SecretValue::new)
} else {
self.read_with_policy(primary, policy).await?
};
let Some(value) = value else {
return Ok(PythonSecretRead::Value(None));
};
if matches!(policy, ReadPolicy::Python) && value.expose().is_empty() {
return Ok(PythonSecretRead::Value(None));
}
Ok(PythonSecretRead::PrimaryJson(value))
}
}
}
async fn read_with_policy(
&self,
name: &str,
policy: ReadPolicy,
) -> Result<Option<SecretValue>, Error> {
match (
Self::read_with_client(&self.client, name, policy).await,
policy,
) {
(Err(Error::Read(_) | Error::MissingString | Error::Timeout), ReadPolicy::Python) => {
Ok(None)
}
(result, _) => result,
}
}
pub async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
Self::async_read_secret_with_client(&self.client, name).await
}
pub(super) async fn async_read_secret_with_client(
client: &Client,
name: &str,
) -> Result<Option<SecretValue>, Error> {
Self::read_with_client(client, name, ReadPolicy::Native).await
}
async fn read_with_client(
client: &Client,
name: &str,
policy: ReadPolicy,
) -> Result<Option<SecretValue>, Error> {
let request = client.get_secret_value().secret_id(name);
let response = match policy {
ReadPolicy::Native => request.send().await,
ReadPolicy::Python => {
request
.customize()
.config_override(
aws_sdk_secretsmanager::config::Builder::new()
.retry_config(RetryConfig::disabled()),
)
.send()
.await
}
};
match response {
Ok(response) => response
.secret_string
.map(SecretValue::new)
.map(Some)
.ok_or(Error::MissingString),
Err(error)
if matches!(
&error,
aws_sdk_secretsmanager::error::SdkError::TimeoutError(_)
) || matches!(&error, aws_sdk_secretsmanager::error::SdkError::DispatchFailure(failure) if failure.is_timeout()) =>
{
Err(Error::Timeout)
}
Err(error)
if error
.as_service_error()
.is_some_and(|error| error.is_resource_not_found_exception()) =>
{
Ok(None)
}
Err(error) => Err(Error::Read(Box::new(error))),
}
}
}
impl BaseSecretManager for AwsSecretsManagerV2 {
type Error = Error;
type Context = AwsOperationContext;
async fn async_read_secret(
&self,
name: &str,
context: &Self::Context,
) -> Result<Option<SecretValue>, Error> {
let client = self.client_for_context(context)?;
Self::async_read_secret_with_client(&client, name).await
}
}
pub fn is_bootstrap_key(name: &str) -> bool {
matches!(
name,
AWS_ACCESS_KEY_ID
| AWS_SECRET_ACCESS_KEY
| AWS_REGION_NAME
| AWS_REGION
| AWS_BEDROCK_RUNTIME_ENDPOINT
)
}
fn resolve_payload(payload: PythonSecretRead, name: &str) -> Result<Option<Secret>, Error> {
match payload {
PythonSecretRead::Value(value) => Ok(value),
PythonSecretRead::PrimaryJson(document) => {
let object: Value =
serde_json::from_str(document.expose()).map_err(|_| Error::PrimarySecret)?;
let object = object.as_object().ok_or(Error::PrimarySecret)?;
Ok(object.get(name).cloned().map(Secret::from_json))
}
}
}

View file

@ -0,0 +1,329 @@
use super::*;
impl AwsSecretsManagerV2 {
pub async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
description: Option<&str>,
) -> Result<CreateSecretOutput, Error> {
self.async_write_secret_with_client_and_tags(&self.client, name, value, description, None)
.await
}
pub(super) async fn async_write_secret_with_client_and_tags(
&self,
client: &Client,
name: &str,
value: &SecretValue,
description: Option<&str>,
tags: Option<&BTreeMap<String, String>>,
) -> Result<CreateSecretOutput, Error> {
let tags = self.write_tags(tags);
let request = client
.create_secret()
.name(name)
.secret_string(value.expose())
.set_description(description.filter(|v| !v.is_empty()).map(str::to_owned))
.set_kms_key_id(self.write_kms_key_id())
.set_tags(tags.clone());
let response = match request.send().await {
Ok(response) => response,
Err(error) => self
.restore_and_update_secret(client, name, value, description, tags)
.await?
.ok_or_else(|| Error::Create(Box::new(error)))?,
};
if let Some(regions) = &self.write_settings.replica_regions
&& !regions.is_empty()
&& self
.async_replicate_secret_with_client(client, name, regions)
.await
.is_err()
{
litellm_tracing::warn!("secret created but replication failed");
}
Ok(response)
}
async fn restore_and_update_secret(
&self,
client: &Client,
name: &str,
value: &SecretValue,
description: Option<&str>,
tags: Option<Vec<Tag>>,
) -> Result<Option<CreateSecretOutput>, Error> {
let scheduled = client
.describe_secret()
.secret_id(name)
.send()
.await
.is_ok_and(|response| response.deleted_date().is_some());
if !scheduled {
return Ok(None);
}
client
.restore_secret()
.secret_id(name)
.send()
.await
.map_err(|error| Error::Restore(Box::new(error)))?;
match self
.update_restored_secret(client, name, value, description, tags)
.await
{
Ok(response) => Ok(Some(response)),
Err(error) => {
self.async_delete_secret_with_client(client, name, Some(7))
.await?;
Err(error)
}
}
}
fn write_kms_key_id(&self) -> Option<String> {
self.write_settings
.kms_key_id
.clone()
.filter(|value| !value.is_empty())
}
fn write_tags(&self, tags: Option<&BTreeMap<String, String>>) -> Option<Vec<Tag>> {
tags.or(self.write_settings.tags.as_ref()).map(|tags| {
tags.iter()
.map(|(key, value)| Tag::builder().key(key).value(value).build())
.collect()
})
}
async fn update_restored_secret(
&self,
client: &Client,
name: &str,
value: &SecretValue,
description: Option<&str>,
tags: Option<Vec<Tag>>,
) -> Result<CreateSecretOutput, Error> {
let response = client
.update_secret()
.secret_id(name)
.secret_string(value.expose())
.set_description(
description
.filter(|value| !value.is_empty())
.map(str::to_owned),
)
.set_kms_key_id(self.write_kms_key_id())
.send()
.await
.map_err(|error| Error::Update(Box::new(error)))?;
if let Some(tags) = tags {
client
.tag_resource()
.secret_id(name)
.set_tags(Some(tags))
.send()
.await
.map_err(|error| Error::Tag(Box::new(error)))?;
}
Ok(CreateSecretOutput::builder()
.set_arn(response.arn)
.set_name(response.name)
.set_version_id(response.version_id)
.build())
}
pub async fn async_replicate_secret(
&self,
name: &str,
regions: &[String],
) -> Result<Option<ReplicateSecretToRegionsOutput>, Error> {
self.async_replicate_secret_with_client(&self.client, name, regions)
.await
}
pub(super) async fn async_replicate_secret_with_client(
&self,
client: &Client,
name: &str,
regions: &[String],
) -> Result<Option<ReplicateSecretToRegionsOutput>, Error> {
if regions.is_empty() {
return Ok(None);
}
client
.replicate_secret_to_regions()
.secret_id(name)
.set_add_replica_regions(Some(
regions
.iter()
.map(|region| ReplicaRegionType::builder().region(region).build())
.collect(),
))
.send()
.await
.map(Some)
.map_err(|error| Error::Replicate(Box::new(error)))
}
pub async fn async_put_secret_value(
&self,
name: &str,
value: &SecretValue,
) -> Result<PutSecretValueOutput, Error> {
self.async_put_secret_value_with_client(&self.client, name, value)
.await
}
pub(super) async fn async_put_secret_value_with_client(
&self,
client: &Client,
name: &str,
value: &SecretValue,
) -> Result<PutSecretValueOutput, Error> {
client
.put_secret_value()
.secret_id(name)
.secret_string(value.expose())
.send()
.await
.map_err(|error| Error::Put(Box::new(error)))
}
pub async fn async_delete_secret(
&self,
name: &str,
recovery_window_in_days: Option<u32>,
) -> Result<DeleteSecretOutput, Error> {
self.async_delete_secret_with_client(&self.client, name, recovery_window_in_days)
.await
}
pub async fn async_delete_secret_with_context(
&self,
name: &str,
recovery_window_in_days: Option<u32>,
context: &AwsOperationContext,
) -> Result<DeleteSecretOutput, Error> {
let client = self.client_for_context(context)?;
self.async_delete_secret_with_client(&client, name, recovery_window_in_days)
.await
}
pub(super) async fn async_delete_secret_with_client(
&self,
client: &Client,
name: &str,
recovery_window_in_days: Option<u32>,
) -> Result<DeleteSecretOutput, Error> {
client
.delete_secret()
.secret_id(name)
.set_recovery_window_in_days(recovery_window_in_days.map(i64::from))
.send()
.await
.map_err(|error| Error::Delete(Box::new(error)))
}
pub async fn async_rotate_secret(
&self,
current_name: &str,
new_name: &str,
value: &SecretValue,
) -> Result<RotationResponse, RotationError<RotationResponse, Error>> {
self.async_rotate_secret_with_context(
current_name,
new_name,
value,
&AwsOperationContext::default(),
)
.await
}
pub async fn async_rotate_secret_with_context(
&self,
current_name: &str,
new_name: &str,
value: &SecretValue,
context: &AwsOperationContext,
) -> Result<RotationResponse, RotationError<RotationResponse, Error>> {
if current_name == new_name {
return self
.async_write_replacement(current_name, new_name, value, context)
.await
.map_err(RotationError::Write);
}
async_rotate_secret(self, current_name, new_name, value, context).await
}
}
impl SecretWriter for AwsSecretsManagerV2 {
type WriteResponse = CreateSecretOutput;
async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
context: &SecretWriteContext<Self::Context>,
) -> Result<CreateSecretOutput, Error> {
let client = self.client_for_context(&context.operation)?;
self.async_write_secret_with_client_and_tags(
&client,
name,
value,
context.description.as_deref(),
(!context.tags.is_empty()).then_some(&context.tags),
)
.await
}
}
impl SecretDeleter for AwsSecretsManagerV2 {
type DeleteResponse = DeleteSecretOutput;
async fn async_delete_secret(
&self,
name: &str,
context: &Self::Context,
) -> Result<DeleteSecretOutput, Error> {
self.async_delete_secret_with_context(name, Some(7), context)
.await
}
}
impl SecretRotator for AwsSecretsManagerV2 {
type RotationResponse = RotationResponse;
async fn async_read_secret_fresh(
&self,
name: &str,
context: &Self::Context,
) -> Result<Option<SecretValue>, Error> {
BaseSecretManager::async_read_secret(self, name, context).await
}
async fn async_write_replacement(
&self,
current_name: &str,
new_name: &str,
value: &SecretValue,
context: &Self::Context,
) -> Result<RotationResponse, Error> {
if current_name == new_name {
let client = self.client_for_context(context)?;
return self
.async_put_secret_value_with_client(&client, new_name, value)
.await
.map(RotationResponse::Updated);
}
SecretWriter::async_write_secret(
self,
new_name,
value,
&SecretWriteContext::rotated_from(current_name, context.clone()),
)
.await
.map(RotationResponse::Created)
}
}

View file

@ -60,10 +60,13 @@ fn disabled_kms_loader_does_not_require_environment_configuration(#[case] enable
}
#[rstest]
#[case::settings(Some("configured-region"), None)]
#[case::environment(None, Some("environment-region"))]
fn enabled_kms_loader_accepts_either_region_source(
#[case::settings(Some("configured-region"), None, None)]
#[case::region_name(None, Some("AWS_REGION_NAME"), Some("environment-region"))]
#[case::region(None, Some("AWS_REGION"), Some("environment-region"))]
#[case::default_region(None, Some("AWS_DEFAULT_REGION"), Some("environment-region"))]
fn enabled_kms_loader_accepts_supported_region_sources(
#[case] configured_region: Option<&'static str>,
#[case] environment_region_name: Option<&'static str>,
#[case] environment_region: Option<&'static str>,
) {
use std::sync::Arc;
@ -72,7 +75,7 @@ fn enabled_kms_loader_accepts_either_region_source(
..KeyManagementSettings::default()
};
let environment = Arc::new(move |name: &str| {
(name == "AWS_REGION_NAME")
(Some(name) == environment_region_name)
.then(|| environment_region.map(str::to_owned))
.flatten()
});

View file

@ -12,8 +12,8 @@ use aws_sdk_secretsmanager::{
};
use litellm_secrets_aws::{AwsSecretsManagerV2, Error, RotationResponse};
use litellm_secrets_types::{
AwsOperationContext, BaseSecretManager, KeyManagementSettings, SecretOperationContext,
SecretValue, SecretWriteContext,
AwsOperationContext, BaseSecretManager, KeyManagementSettings, Secret, SecretDeleter,
SecretValue, SecretWriteContext, SecretWriter,
};
use rstest::{fixture, rstest};
use serde_json::json;
@ -22,459 +22,13 @@ use wiremock::{
matchers::{body_partial_json, header},
};
fn manager(server: &MockServer, settings: KeyManagementSettings) -> AwsSecretsManagerV2 {
let client = Client::from_conf(
aws_sdk_secretsmanager::Config::builder()
.behavior_version(BehaviorVersion::latest())
.region(Region::new("us-east-1"))
.credentials_provider(Credentials::new("test", "test", None, None, "test"))
.endpoint_url(server.uri())
.retry_config(RetryConfig::disabled())
.build(),
);
AwsSecretsManagerV2::new(client, (&settings).into())
}
#[path = "secret_manager/support.rs"]
mod support;
use support::*;
fn loaded_manager(server: &MockServer) -> AwsSecretsManagerV2 {
let endpoint_url = server.uri();
let environment: Arc<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
Arc::new(move |name: &str| match name {
"AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint_url.clone()),
"AWS_ACCESS_KEY_ID" => Some("test".into()),
"AWS_SECRET_ACCESS_KEY" => Some("test".into()),
_ => None,
});
AwsSecretsManagerV2::load_aws_secret_manager(
Some(true),
KeyManagementSettings {
aws_region_name: Some("us-east-1".into()),
..Default::default()
},
environment,
)
.unwrap()
.unwrap()
}
#[fixture]
fn default_settings() -> KeyManagementSettings {
KeyManagementSettings::default()
}
#[rstest]
#[case::string_value("KEY", Some("value"))]
#[case::missing_value("missing", None)]
#[case::non_string_value("BOOL", None)]
#[tokio::test]
async fn primary_lookup_preserves_read_semantics(
default_settings: KeyManagementSettings,
#[case] name: &str,
#[case] expected: Option<&str>,
) {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.and(body_partial_json(json!({"SecretId":"primary"})))
.respond_with(
ResponseTemplate::new(200).set_body_json(
json!({"SecretString":json!({"KEY":"value", "BOOL":true}).to_string()}),
),
)
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, default_settings);
assert_eq!(
manager
.read_secret_for_resolver(name, Some("primary"), &|_: &str| None)
.await
.unwrap()
.and_then(|v| v.as_str().map(str::to_owned))
.as_deref(),
expected
);
}
#[rstest]
#[case::access_key("AWS_ACCESS_KEY_ID")]
#[case::secret_access_key("AWS_SECRET_ACCESS_KEY")]
#[case::region_name("AWS_REGION_NAME")]
#[case::region("AWS_REGION")]
#[case::bedrock_endpoint("AWS_BEDROCK_RUNTIME_ENDPOINT")]
#[tokio::test]
async fn bootstrap_keys_bypass_primary_lookup(
default_settings: KeyManagementSettings,
#[case] name: &str,
) {
let server = MockServer::start().await;
let manager = manager(&server, default_settings);
assert_eq!(
manager
.read_secret_for_resolver(name, Some("primary"), &|_: &str| Some("bootstrap".into()))
.await
.unwrap()
.unwrap()
.as_str()
.unwrap(),
"bootstrap"
);
}
#[rstest]
#[tokio::test]
async fn failed_read_returns_none_but_invalid_primary_json_is_an_error(
default_settings: KeyManagementSettings,
) {
let server = MockServer::start().await;
Mock::given(body_partial_json(json!({"SecretId":"missing"})))
.respond_with(
ResponseTemplate::new(400).set_body_json(json!({"__type":"ResourceNotFoundException"})),
)
.mount(&server)
.await;
Mock::given(body_partial_json(json!({"SecretId":"invalid"})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString":"not-json"})))
.mount(&server)
.await;
Mock::given(body_partial_json(json!({"SecretId":"no-string"})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"no-string"})))
.mount(&server)
.await;
let manager = manager(&server, default_settings);
assert!(
manager
.async_read_secret("missing")
.await
.unwrap()
.is_none()
);
assert!(matches!(
manager
.read_secret_for_resolver("KEY", Some("invalid"), &|_: &str| None)
.await,
Err(Error::PrimarySecret)
));
assert!(matches!(
manager.async_read_secret("no-string").await,
Err(Error::MissingString)
));
}
#[rstest]
#[tokio::test]
async fn same_name_rotation_uses_put_and_returns_its_response(
default_settings: KeyManagementSettings,
) {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.PutSecretValue"))
.and(body_partial_json(
json!({"SecretId":"key", "SecretString":"replacement"}),
))
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"Name":"key", "VersionId":"version"})),
)
.expect(1)
.mount(&server)
.await;
let response = manager(&server, default_settings)
.async_rotate_secret("key", "key", &SecretValue::new("replacement"))
.await
.unwrap();
match response {
RotationResponse::Updated(output) => assert_eq!(output.version_id(), Some("version")),
_ => panic!("rotation created a second secret"),
}
assert_eq!(server.received_requests().await.unwrap().len(), 1);
}
#[rstest]
#[tokio::test]
async fn renamed_rotation_reads_creates_verifies_then_deletes(
default_settings: KeyManagementSettings,
) {
let server = MockServer::start().await;
let step = AtomicUsize::new(0);
Mock::given(wiremock::matchers::method("POST"))
.respond_with(move |request: &wiremock::Request| {
let body: serde_json::Value = request.body_json().unwrap();
let action = request
.headers
.get("x-amz-target")
.unwrap()
.to_str()
.unwrap();
match step.fetch_add(1, Ordering::SeqCst) {
0 => {
assert_eq!(action, "secretsmanager.GetSecretValue");
assert_eq!(body["SecretId"], "old");
ResponseTemplate::new(200).set_body_json(json!({"SecretString":"old-value"}))
}
1 => {
assert_eq!(action, "secretsmanager.CreateSecret");
assert_eq!(body["Name"], "new");
assert_eq!(body["Description"], "Rotated from old");
assert_eq!(body["SecretString"], "replacement");
ResponseTemplate::new(200).set_body_json(json!({"Name":"new"}))
}
2 => {
assert_eq!(action, "secretsmanager.GetSecretValue");
assert_eq!(body["SecretId"], "new");
ResponseTemplate::new(200).set_body_json(json!({"SecretString":"replacement"}))
}
3 => {
assert_eq!(action, "secretsmanager.DeleteSecret");
assert_eq!(body["SecretId"], "old");
assert_eq!(body["RecoveryWindowInDays"], 7);
ResponseTemplate::new(200).set_body_json(json!({"Name":"old"}))
}
_ => panic!("unexpected request"),
}
})
.expect(4)
.mount(&server)
.await;
assert!(matches!(
manager(&server, default_settings)
.async_rotate_secret("old", "new", &SecretValue::new("replacement"))
.await
.unwrap(),
RotationResponse::Created(_)
));
}
#[rstest]
#[tokio::test]
async fn creation_passes_tags_and_kms_and_survives_replication_failure() {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.CreateSecret"))
.and(body_partial_json(json!({"Name":"key", "SecretString":"value", "KmsKeyId":"kms-key", "Tags":[{"Key":"stage", "Value":"test"}]})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"key"}))).expect(1).mount(&server).await;
Mock::given(header(
"x-amz-target",
"secretsmanager.ReplicateSecretToRegions",
))
.and(body_partial_json(
json!({"SecretId":"key", "AddReplicaRegions":[{"Region":"replica-region"}]}),
))
.respond_with(
ResponseTemplate::new(400).set_body_json(json!({"__type":"InvalidRequestException"})),
)
.expect(1)
.mount(&server)
.await;
let settings = KeyManagementSettings {
kms_key_id: Some("kms-key".into()),
tags: Some(std::collections::BTreeMap::from([(
"stage".into(),
"test".into(),
)])),
replica_regions: Some(vec!["replica-region".into()]),
..Default::default()
};
let manager = manager(&server, settings);
assert_eq!(
manager
.async_write_secret("key", &SecretValue::new("value"), None)
.await
.unwrap()
.name(),
Some("key")
);
assert!(
manager
.async_replicate_secret("key", &[])
.await
.unwrap()
.is_none()
);
}
#[rstest]
#[tokio::test]
async fn trait_write_uses_typed_write_context(default_settings: KeyManagementSettings) {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.CreateSecret"))
.and(body_partial_json(json!({
"Name": "key",
"SecretString": "value",
"Description": "created by caller",
"Tags": [{"Key": "stage", "Value": "test"}],
})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name": "key"})))
.expect(1)
.mount(&server)
.await;
let context = SecretWriteContext {
description: Some("created by caller".into()),
tags: std::collections::BTreeMap::from([("stage".into(), "test".into())]),
..Default::default()
};
let response = BaseSecretManager::async_write_secret(
&manager(&server, default_settings),
"key",
&SecretValue::new("value"),
&context,
)
.await
.unwrap();
assert_eq!(response.name(), Some("key"));
}
#[rstest]
#[tokio::test]
async fn trait_delete_accepts_an_unspecified_recovery_window(
default_settings: KeyManagementSettings,
) {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.DeleteSecret"))
.and(body_partial_json(json!({"SecretId": "key"})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name": "key"})))
.expect(1)
.mount(&server)
.await;
let response = BaseSecretManager::async_delete_secret(
&manager(&server, default_settings),
"key",
None,
&SecretOperationContext::default(),
)
.await
.unwrap();
assert_eq!(response.name(), Some("key"));
}
#[rstest]
#[tokio::test]
async fn trait_read_uses_the_aws_region_from_its_operation_context() {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.respond_with(|request: &wiremock::Request| {
let authorization = request
.headers
.get("authorization")
.unwrap()
.to_str()
.unwrap();
assert!(authorization.contains("/us-west-2/secretsmanager/aws4_request"));
ResponseTemplate::new(200).set_body_json(json!({"SecretString": "value"}))
})
.expect(1)
.mount(&server)
.await;
let context = SecretOperationContext::Aws(AwsOperationContext {
region_name: Some("us-west-2".into()),
..Default::default()
});
let value = BaseSecretManager::async_read_secret(&loaded_manager(&server), "key", &context)
.await
.unwrap();
assert_eq!(value.unwrap().expose(), "value");
}
#[rstest]
#[tokio::test]
async fn trait_read_applies_the_aws_operation_timeout() {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(1))
.set_body_json(json!({"SecretString": "late"})),
)
.expect(1)
.mount(&server)
.await;
let context = SecretOperationContext::Aws(AwsOperationContext {
timeout: Some(Duration::from_millis(30)),
..Default::default()
});
assert!(matches!(
BaseSecretManager::async_read_secret(&loaded_manager(&server), "key", &context).await,
Err(Error::Timeout)
));
}
#[rstest]
#[tokio::test]
async fn credential_failures_are_not_swallowed_as_missing_secrets() {
use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future};
#[derive(Debug)]
struct FailedCredentials;
impl ProvideCredentials for FailedCredentials {
fn provide_credentials<'a>(&'a self) -> future::ProvideCredentials<'a>
where
Self: 'a,
{
future::ProvideCredentials::ready(Err(CredentialsError::provider_error(
"private-auth-detail",
)))
}
}
let server = MockServer::start().await;
let config = aws_sdk_secretsmanager::Config::builder()
.behavior_version(BehaviorVersion::latest())
.region(Region::new("us-east-1"))
.credentials_provider(FailedCredentials)
.endpoint_url(server.uri())
.retry_config(RetryConfig::disabled())
.build();
let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default());
let error = manager.async_read_secret("key").await.unwrap_err();
assert!(!format!("{error:?}").contains("private-auth-detail"));
assert!(matches!(error, Error::Read(_)));
assert!(server.received_requests().await.unwrap().is_empty());
}
#[rstest]
#[tokio::test]
async fn read_timeout_is_an_error_and_cannot_be_mistaken_for_missing() {
let server = MockServer::start().await;
Mock::given(wiremock::matchers::method("POST"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(1))
.set_body_json(json!({"SecretString":"late"})),
)
.mount(&server)
.await;
let config = aws_sdk_secretsmanager::Config::builder()
.behavior_version(BehaviorVersion::latest())
.region(Region::new("us-east-1"))
.credentials_provider(Credentials::new("test", "test", None, None, "test"))
.endpoint_url(server.uri())
.retry_config(RetryConfig::disabled())
.timeout_config(
aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder()
.operation_timeout(Duration::from_millis(30))
.build(),
)
.build();
let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default());
assert!(matches!(
manager.async_read_secret("key").await,
Err(Error::Timeout)
));
}
#[rstest]
#[case::denied(400, "AccessDeniedException")]
#[case::throttled(400, "ThrottlingException")]
#[case::unavailable(503, "ServiceUnavailableException")]
#[tokio::test]
async fn service_failures_remain_errors(
default_settings: KeyManagementSettings,
#[case] status: u16,
#[case] code: &str,
) {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.respond_with(ResponseTemplate::new(status).set_body_json(json!({"__type":code})))
.expect(1)
.mount(&server)
.await;
assert!(matches!(
manager(&server, default_settings)
.async_read_secret("key")
.await,
Err(Error::Read(_))
));
}
#[path = "secret_manager/configuration.rs"]
mod configuration;
#[path = "secret_manager/reads.rs"]
mod reads;
#[path = "secret_manager/writes.rs"]
mod writes;

View file

@ -0,0 +1,256 @@
use super::*;
#[rstest]
#[tokio::test]
async fn credential_failures_are_not_swallowed_as_missing_secrets() {
use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future};
#[derive(Debug)]
struct FailedCredentials;
impl ProvideCredentials for FailedCredentials {
fn provide_credentials<'a>(&'a self) -> future::ProvideCredentials<'a>
where
Self: 'a,
{
future::ProvideCredentials::ready(Err(CredentialsError::provider_error(
"private-auth-detail",
)))
}
}
let server = MockServer::start().await;
let config = client_builder(&server)
.credentials_provider(FailedCredentials)
.build();
let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default());
let error = manager.async_read_secret("key").await.unwrap_err();
assert!(!format!("{error:?}").contains("private-auth-detail"));
assert!(matches!(error, Error::Read(_)));
assert!(server.received_requests().await.unwrap().is_empty());
}
#[rstest]
#[case::environment(false)]
#[case::operation_override(true)]
#[tokio::test]
async fn endpoint_overrides_replace_the_service_and_override_the_region(
#[case] override_context: bool,
) {
let configured = MockServer::start().await;
let explicit = MockServer::start().await;
let target = if override_context {
&explicit
} else {
&configured
};
Mock::given(wiremock::matchers::path_regex("^/secretsmanager/?$"))
.and(header("x-amz-target", "secretsmanager.GetSecretValue"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString":"value"})))
.expect(1)
.mount(target)
.await;
let endpoint = format!("{}/bedrock-runtime", configured.uri());
let manager = AwsSecretsManagerV2::load_aws_secret_manager(
Some(true),
KeyManagementSettings {
aws_region_name: Some("cn-north-1".into()),
..Default::default()
},
Arc::new(move |name: &str| match name {
"AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()),
"AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("test".into()),
_ => None,
}),
)
.unwrap()
.unwrap();
let context = AwsOperationContext {
bedrock_runtime_endpoint: override_context
.then(|| format!("{}/bedrock-runtime", explicit.uri())),
..Default::default()
};
assert_eq!(
BaseSecretManager::async_read_secret(&manager, "key", &context)
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
assert!(
if override_context {
configured
} else {
explicit
}
.received_requests()
.await
.unwrap()
.is_empty()
);
}
#[rstest]
#[case::unset(None)]
#[case::disabled(Some(false))]
fn disabled_secret_manager_loader_does_not_require_environment(#[case] enabled: Option<bool>) {
assert!(
AwsSecretsManagerV2::load_aws_secret_manager(
enabled,
Default::default(),
Arc::new(|_: &str| panic!("disabled loader consulted the environment"))
)
.unwrap()
.is_none()
);
}
#[rstest]
#[case::role(None, None)]
#[case::cross_account(Some("external-id"), None)]
#[case::web_identity(None, Some("identity-token"))]
#[tokio::test]
async fn configured_sts_credentials_sign_the_secret_request(
#[case] external_id: Option<&str>,
#[case] identity: Option<&str>,
) {
use aws_sdk_secretsmanager::primitives::{DateTime, DateTimeFormat};
let server = MockServer::start().await;
let expiry = DateTime::from(std::time::SystemTime::now() + Duration::from_secs(3600))
.fmt(DateTimeFormat::DateTime)
.unwrap();
let action = if identity.is_some() {
"AssumeRoleWithWebIdentity"
} else {
"AssumeRole"
};
let expected_external = external_id.map(str::to_owned);
let expected_identity = identity.map(str::to_owned);
Mock::given(wiremock::matchers::body_string_contains(format!("Action={action}")))
.respond_with(move |request: &wiremock::Request| {
let body = std::str::from_utf8(&request.body).unwrap();
assert!(body.contains("RoleArn=test-role"), "{body}");
assert!(body.contains("RoleSessionName=parity-session"), "{body}");
if let Some(value) = &expected_external { assert!(body.contains(&format!("ExternalId={value}"))); }
if let Some(value) = &expected_identity { assert!(body.contains(&format!("WebIdentityToken={value}"))); }
ResponseTemplate::new(200).set_body_string(format!(
"<{action}Response><{action}Result><Credentials><AccessKeyId>assumed-key</AccessKeyId>\
<SecretAccessKey>assumed-secret</SecretAccessKey><SessionToken>session-token</SessionToken>\
<Expiration>{expiry}</Expiration></Credentials></{action}Result></{action}Response>"))
}).expect(1).mount(&server).await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.and(header("x-amz-security-token", "session-token"))
.respond_with(|request: &wiremock::Request| {
assert!(
request.headers["authorization"]
.to_str()
.unwrap()
.contains("Credential=assumed-key/")
);
ResponseTemplate::new(200).set_body_json(json!({"SecretString":"value"}))
})
.expect(1)
.mount(&server)
.await;
let endpoint = server.uri();
let manager = AwsSecretsManagerV2::load_aws_secret_manager(
Some(true),
KeyManagementSettings {
aws_region_name: Some("us-east-1".into()),
aws_role_name: Some("test-role".into()),
aws_session_name: Some("parity-session".into()),
aws_external_id: external_id.map(SecretValue::new),
aws_web_identity_token: identity.map(SecretValue::new),
aws_sts_endpoint: Some(endpoint.clone()),
..Default::default()
},
Arc::new(move |name: &str| match name {
"AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()),
"AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("source-key".into()),
_ => None,
}),
)
.unwrap()
.unwrap();
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[tokio::test]
async fn configured_profile_credentials_override_static_environment_credentials() {
const CHILD_ENDPOINT: &str = "LITELLM_SECRETS_PROFILE_TEST_ENDPOINT";
if let Ok(endpoint) = std::env::var(CHILD_ENDPOINT) {
let manager = AwsSecretsManagerV2::load_aws_secret_manager(
Some(true),
KeyManagementSettings {
aws_region_name: Some("us-east-1".into()),
aws_profile_name: Some("parity".into()),
..Default::default()
},
Arc::new(move |name: &str| match name {
"AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()),
"AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("wrong-static-key".into()),
_ => None,
}),
)
.unwrap()
.unwrap();
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"profile-value"
);
return;
}
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.and(header("x-amz-security-token", "profile-session"))
.respond_with(|request: &wiremock::Request| {
assert!(
request.headers["authorization"]
.to_str()
.unwrap()
.contains("Credential=profile-key/")
);
ResponseTemplate::new(200).set_body_json(json!({"SecretString":"profile-value"}))
})
.expect(1)
.mount(&server)
.await;
let directory = tempfile::tempdir().unwrap();
let credentials = directory.path().join("credentials");
let config = directory.path().join("config");
std::fs::write(&credentials, "[parity]\naws_access_key_id=profile-key\naws_secret_access_key=profile-secret\naws_session_token=profile-session\n").unwrap();
std::fs::write(&config, "").unwrap();
let endpoint = server.uri();
let result = tokio::task::spawn_blocking(move || {
std::process::Command::new(std::env::current_exe().unwrap())
.args([
"--exact",
"configuration::configured_profile_credentials_override_static_environment_credentials",
"--nocapture",
])
.env(CHILD_ENDPOINT, endpoint)
.env("AWS_SHARED_CREDENTIALS_FILE", credentials)
.env("AWS_CONFIG_FILE", config)
.output()
.unwrap()
})
.await
.unwrap();
assert!(
result.status.success(),
"{}\n{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
}

View file

@ -0,0 +1,198 @@
use super::*;
#[rstest]
#[case::string_value("KEY", Some(Secret::String(SecretValue::new("value"))))]
#[case::missing_value("missing", None)]
#[case::non_string_value("BOOL", Some(Secret::Bool(true)))]
#[tokio::test]
async fn primary_lookup_preserves_read_semantics(
default_settings: KeyManagementSettings,
#[case] name: &str,
#[case] expected: Option<Secret>,
) {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.and(body_partial_json(json!({"SecretId":"primary"})))
.respond_with(
ResponseTemplate::new(200).set_body_json(
json!({"SecretString":json!({"KEY":"value", "BOOL":true}).to_string()}),
),
)
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, default_settings);
assert_eq!(
manager
.read_secret_for_resolver(name, Some("primary"), &|_: &str| None)
.await
.unwrap(),
expected
);
}
#[rstest]
#[case::access_key("AWS_ACCESS_KEY_ID")]
#[case::secret_access_key("AWS_SECRET_ACCESS_KEY")]
#[case::region_name("AWS_REGION_NAME")]
#[case::region("AWS_REGION")]
#[case::bedrock_endpoint("AWS_BEDROCK_RUNTIME_ENDPOINT")]
#[tokio::test]
async fn bootstrap_keys_bypass_primary_lookup(
default_settings: KeyManagementSettings,
#[case] name: &str,
) {
let server = MockServer::start().await;
let manager = manager(&server, default_settings);
assert_eq!(
manager
.read_secret_for_resolver(name, Some("primary"), &|_: &str| Some("bootstrap".into()))
.await
.unwrap()
.unwrap()
.as_str()
.unwrap(),
"bootstrap"
);
}
#[rstest]
#[tokio::test]
async fn failed_read_returns_none_but_invalid_primary_json_is_an_error(
default_settings: KeyManagementSettings,
) {
let server = MockServer::start().await;
Mock::given(body_partial_json(json!({"SecretId":"missing"})))
.respond_with(
ResponseTemplate::new(400).set_body_json(json!({"__type":"ResourceNotFoundException"})),
)
.mount(&server)
.await;
Mock::given(body_partial_json(json!({"SecretId":"invalid"})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString":"not-json"})))
.mount(&server)
.await;
Mock::given(body_partial_json(json!({"SecretId":"no-string"})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"no-string"})))
.mount(&server)
.await;
let manager = manager(&server, default_settings);
assert!(
manager
.async_read_secret("missing")
.await
.unwrap()
.is_none()
);
assert!(matches!(
manager
.read_secret_for_resolver("KEY", Some("invalid"), &|_: &str| None)
.await,
Err(Error::PrimarySecret)
));
assert!(matches!(
manager.async_read_secret("no-string").await,
Err(Error::MissingString)
));
}
#[rstest]
#[tokio::test]
async fn trait_read_uses_the_aws_region_from_its_operation_context() {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.respond_with(|request: &wiremock::Request| {
let authorization = request
.headers
.get("authorization")
.unwrap()
.to_str()
.unwrap();
assert!(authorization.contains("/us-west-2/secretsmanager/aws4_request"));
ResponseTemplate::new(200).set_body_json(json!({"SecretString": "value"}))
})
.expect(1)
.mount(&server)
.await;
let context = AwsOperationContext {
region_name: Some("us-west-2".into()),
..Default::default()
};
let value = BaseSecretManager::async_read_secret(&loaded_manager(&server), "key", &context)
.await
.unwrap();
assert_eq!(value.unwrap().expose(), "value");
}
#[rstest]
#[tokio::test]
async fn trait_read_applies_the_aws_operation_timeout() {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(1))
.set_body_json(json!({"SecretString": "late"})),
)
.expect(1)
.mount(&server)
.await;
let context = AwsOperationContext {
timeout: Some(Duration::from_millis(30)),
..Default::default()
};
assert!(matches!(
BaseSecretManager::async_read_secret(&loaded_manager(&server), "key", &context).await,
Err(Error::Timeout)
));
}
#[rstest]
#[tokio::test]
async fn read_timeout_is_an_error_and_cannot_be_mistaken_for_missing() {
let server = MockServer::start().await;
Mock::given(wiremock::matchers::method("POST"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(1))
.set_body_json(json!({"SecretString":"late"})),
)
.mount(&server)
.await;
let config = client_builder(&server)
.timeout_config(
aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder()
.operation_timeout(Duration::from_millis(30))
.build(),
)
.build();
let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default());
assert!(matches!(
manager.async_read_secret("key").await,
Err(Error::Timeout)
));
}
#[rstest]
#[case::denied(400, "AccessDeniedException")]
#[case::throttled(400, "ThrottlingException")]
#[case::unavailable(503, "ServiceUnavailableException")]
#[tokio::test]
async fn service_failures_remain_errors(
default_settings: KeyManagementSettings,
#[case] status: u16,
#[case] code: &str,
) {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
.respond_with(ResponseTemplate::new(status).set_body_json(json!({"__type":code})))
.expect(1)
.mount(&server)
.await;
assert!(matches!(
manager(&server, default_settings)
.async_read_secret("key")
.await,
Err(Error::Read(_))
));
}

View file

@ -0,0 +1,80 @@
use super::*;
pub(super) fn manager(server: &MockServer, settings: KeyManagementSettings) -> AwsSecretsManagerV2 {
let client = Client::from_conf(client_builder(server).build());
AwsSecretsManagerV2::new(client, (&settings).into())
}
pub(super) fn loaded_manager(server: &MockServer) -> AwsSecretsManagerV2 {
let endpoint_url = server.uri();
let environment: Arc<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
Arc::new(move |name: &str| match name {
"AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint_url.clone()),
"AWS_ACCESS_KEY_ID" => Some("test".into()),
"AWS_SECRET_ACCESS_KEY" => Some("test".into()),
_ => None,
});
AwsSecretsManagerV2::load_aws_secret_manager(
Some(true),
KeyManagementSettings {
aws_region_name: Some("us-east-1".into()),
..Default::default()
},
environment,
)
.unwrap()
.unwrap()
}
#[fixture]
pub(super) fn default_settings() -> KeyManagementSettings {
KeyManagementSettings::default()
}
pub(super) async fn scripted_actions(server: &MockServer, actions: Vec<Action>) {
let count = actions.len() as u64;
let step = AtomicUsize::new(0);
Mock::given(wiremock::matchers::method("POST"))
.respond_with(move |request: &wiremock::Request| {
let Action {
operation: action,
request: expected,
status,
response,
} = &actions[step.fetch_add(1, Ordering::SeqCst)];
assert_eq!(
request.headers["x-amz-target"],
format!("secretsmanager.{action}")
);
let body: serde_json::Value = request.body_json().unwrap();
let actual = serde_json::Value::Object(
body.as_object()
.unwrap()
.iter()
.filter(|(key, _)| key.as_str() != "ClientRequestToken")
.map(|(key, value)| (key.clone(), value.clone()))
.collect(),
);
assert_eq!(&actual, expected);
ResponseTemplate::new(*status).set_body_json(response)
})
.expect(count)
.mount(server)
.await;
}
pub(super) struct Action {
pub(super) operation: &'static str,
pub(super) request: serde_json::Value,
pub(super) status: u16,
pub(super) response: serde_json::Value,
}
pub(super) fn client_builder(server: &MockServer) -> aws_sdk_secretsmanager::config::Builder {
aws_sdk_secretsmanager::Config::builder()
.behavior_version(BehaviorVersion::latest())
.region(Region::new("us-east-1"))
.credentials_provider(Credentials::new("test", "test", None, None, "test"))
.endpoint_url(server.uri())
.retry_config(RetryConfig::disabled())
}

View file

@ -0,0 +1,615 @@
use super::*;
#[rstest]
#[tokio::test]
async fn same_name_rotation_uses_put_and_returns_its_response(
default_settings: KeyManagementSettings,
) {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.PutSecretValue"))
.and(body_partial_json(
json!({"SecretId":"key", "SecretString":"replacement"}),
))
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"Name":"key", "VersionId":"version"})),
)
.expect(1)
.mount(&server)
.await;
let response = manager(&server, default_settings)
.async_rotate_secret("key", "key", &SecretValue::new("replacement"))
.await
.unwrap();
match response {
RotationResponse::Updated(output) => assert_eq!(output.version_id(), Some("version")),
_ => panic!("rotation created a second secret"),
}
assert_eq!(server.received_requests().await.unwrap().len(), 1);
}
#[rstest]
#[tokio::test]
async fn renamed_rotation_reads_creates_verifies_then_deletes(
default_settings: KeyManagementSettings,
) {
let server = MockServer::start().await;
scripted_actions(
&server,
vec![
Action {
operation: "GetSecretValue",
request: json!({"SecretId":"old"}),
status: 200,
response: json!({"SecretString":"old-value"}),
},
Action {
operation: "CreateSecret",
request: json!({"Name":"new", "Description":"Rotated from old", "SecretString":"replacement"}),
status: 200,
response: json!({"Name":"new"}),
},
Action {
operation: "GetSecretValue",
request: json!({"SecretId":"new"}),
status: 200,
response: json!({"SecretString":"replacement"}),
},
Action {
operation: "DeleteSecret",
request: json!({"SecretId":"old", "RecoveryWindowInDays":7}),
status: 200,
response: json!({"Name":"old"}),
},
],
).await;
assert!(matches!(
manager(&server, default_settings)
.async_rotate_secret("old", "new", &SecretValue::new("replacement"))
.await
.unwrap(),
RotationResponse::Created(_)
));
}
#[rstest]
#[tokio::test]
async fn creation_passes_tags_and_kms_and_survives_replication_failure() {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.CreateSecret"))
.and(body_partial_json(json!({"Name":"key", "SecretString":"value", "KmsKeyId":"kms-key", "Tags":[{"Key":"stage", "Value":"test"}]})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"key"}))).expect(1).mount(&server).await;
Mock::given(header(
"x-amz-target",
"secretsmanager.ReplicateSecretToRegions",
))
.and(body_partial_json(
json!({"SecretId":"key", "AddReplicaRegions":[{"Region":"replica-region"}]}),
))
.respond_with(
ResponseTemplate::new(400).set_body_json(json!({"__type":"InvalidRequestException"})),
)
.expect(1)
.mount(&server)
.await;
let settings = KeyManagementSettings {
kms_key_id: Some("kms-key".into()),
tags: Some(std::collections::BTreeMap::from([(
"stage".into(),
"test".into(),
)])),
replica_regions: Some(vec!["replica-region".into()]),
..Default::default()
};
let manager = manager(&server, settings);
assert_eq!(
manager
.async_write_secret("key", &SecretValue::new("value"), None)
.await
.unwrap()
.name(),
Some("key")
);
assert!(
manager
.async_replicate_secret("key", &[])
.await
.unwrap()
.is_none()
);
}
#[rstest]
#[tokio::test]
async fn trait_write_uses_typed_write_context(default_settings: KeyManagementSettings) {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.CreateSecret"))
.and(body_partial_json(json!({
"Name": "key",
"SecretString": "value",
"Description": "created by caller",
"Tags": [{"Key": "stage", "Value": "test"}],
})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name": "key"})))
.expect(1)
.mount(&server)
.await;
let context = SecretWriteContext {
description: Some("created by caller".into()),
tags: std::collections::BTreeMap::from([("stage".into(), "test".into())]),
..Default::default()
};
let response = SecretWriter::async_write_secret(
&manager(&server, default_settings),
"key",
&SecretValue::new("value"),
&context,
)
.await
.unwrap();
assert_eq!(response.name(), Some("key"));
}
#[rstest]
#[tokio::test]
async fn trait_delete_uses_the_provider_recovery_policy(default_settings: KeyManagementSettings) {
let server = MockServer::start().await;
Mock::given(header("x-amz-target", "secretsmanager.DeleteSecret"))
.and(body_partial_json(json!({"SecretId": "key"})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name": "key"})))
.expect(1)
.mount(&server)
.await;
let response = SecretDeleter::async_delete_secret(
&manager(&server, default_settings),
"key",
&AwsOperationContext::default(),
)
.await
.unwrap();
assert_eq!(response.name(), Some("key"));
}
#[rstest]
#[case::write(false)]
#[case::rotate_back(true)]
#[tokio::test]
async fn recovery_window_alias_is_restored_updated_and_tagged(#[case] rotate: bool) {
let server = MockServer::start().await;
let description = if rotate {
"Rotated from old"
} else {
"description"
};
let write = json!({"Name":"key", "SecretString":"new", "Description":description,
"KmsKeyId":"kms", "Tags":[{"Key":"stage", "Value":"test"}]});
let actions = if rotate {
vec![Action {
operation: "GetSecretValue",
request: json!({"SecretId":"old"}),
status: 200,
response: json!({"SecretString":"old"}),
}]
} else {
vec![]
};
let recovery = vec![
Action {
operation: "CreateSecret",
request: write,
status: 400,
response: json!({"__type":"ResourceExistsException"}),
},
Action {
operation: "DescribeSecret",
request: json!({"SecretId":"key"}),
status: 200,
response: json!({"DeletedDate":1}),
},
Action {
operation: "RestoreSecret",
request: json!({"SecretId":"key"}),
status: 200,
response: json!({"Name":"key"}),
},
Action {
operation: "UpdateSecret",
request: json!({"SecretId":"key", "SecretString":"new", "Description":description,
"KmsKeyId":"kms"}),
status: 200,
response: json!({"ARN":"restored-arn", "Name":"key", "VersionId":"new-version"}),
},
Action {
operation: "TagResource",
request: json!({"SecretId":"key", "Tags":[{"Key":"stage", "Value":"test"}]}),
status: 200,
response: json!({}),
},
];
let verification = if rotate {
vec![
Action {
operation: "GetSecretValue",
request: json!({"SecretId":"key"}),
status: 200,
response: json!({"SecretString":"new"}),
},
Action {
operation: "DeleteSecret",
request: json!({"SecretId":"old", "RecoveryWindowInDays":7}),
status: 200,
response: json!({}),
},
]
} else {
vec![]
};
scripted_actions(
&server,
actions
.into_iter()
.chain(recovery)
.chain(verification)
.collect(),
)
.await;
let manager = manager(
&server,
KeyManagementSettings {
kms_key_id: Some("kms".into()),
tags: Some(std::collections::BTreeMap::from([(
"stage".into(),
"test".into(),
)])),
..Default::default()
},
);
let output = if rotate {
match manager
.async_rotate_secret("old", "key", &SecretValue::new("new"))
.await
.unwrap()
{
RotationResponse::Created(output) => output,
_ => panic!("expected restored alias"),
}
} else {
manager
.async_write_secret("key", &SecretValue::new("new"), Some(description))
.await
.unwrap()
};
assert_eq!(
(output.arn(), output.name(), output.version_id()),
(Some("restored-arn"), Some("key"), Some("new-version"))
);
}
#[rstest]
#[case::live(200, json!({"Name":"key"}))]
#[case::missing(400, json!({"__type":"ResourceNotFoundException"}))]
#[case::denied(400, json!({"__type":"AccessDeniedException"}))]
#[tokio::test]
async fn create_failure_does_not_overwrite_an_alias_without_a_deletion_date(
#[case] status: u16,
#[case] described: serde_json::Value,
) {
let server = MockServer::start().await;
scripted_actions(
&server,
vec![
Action {
operation: "CreateSecret",
request: json!({"Name":"key", "SecretString":"new"}),
status: 400,
response: json!({"__type":"ResourceExistsException"}),
},
Action {
operation: "DescribeSecret",
request: json!({"SecretId":"key"}),
status,
response: described,
},
],
)
.await;
assert!(matches!(
manager(&server, Default::default())
.async_write_secret("key", &SecretValue::new("new"), None)
.await,
Err(Error::Create(_))
));
}
#[derive(Clone, Copy, Debug)]
enum RecoveryFailure {
Restore,
Update,
Tag,
DeleteAfterUpdate,
DeleteAfterTag,
}
#[rstest]
#[case::unconfigured(None)]
#[case::empty(Some(vec![]))]
#[case::configured(Some(vec!["region-a".into(), "region-b".into()]))]
#[tokio::test]
async fn creation_replicates_only_to_configured_regions(#[case] regions: Option<Vec<String>>) {
let server = MockServer::start().await;
let create = vec![Action {
operation: "CreateSecret",
request: json!({"Name":"key", "SecretString":"value", "KmsKeyId":"kms-key"}),
status: 200,
response: json!({"Name":"key", "VersionId":"created"}),
}];
let replicate = regions
.as_ref()
.filter(|regions| !regions.is_empty())
.map(|regions| Action {
operation: "ReplicateSecretToRegions",
request: json!({"SecretId":"key", "AddReplicaRegions":regions.iter()
.map(|region| json!({"Region":region})).collect::<Vec<_>>()}),
status: 200,
response: json!({"ARN":"replica-arn"}),
});
scripted_actions(&server, create.into_iter().chain(replicate).collect()).await;
let environment: Arc<dyn litellm_core_utils::settings::Lookup + Send + Sync> = {
let endpoint = server.uri();
Arc::new(move |name: &str| match name {
"AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()),
"AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("test".into()),
_ => None,
})
};
let manager = AwsSecretsManagerV2::load_aws_secret_manager(
Some(true),
KeyManagementSettings {
aws_region_name: Some("us-east-1".into()),
replica_regions: regions,
kms_key_id: Some("kms-key".into()),
..Default::default()
},
environment,
)
.unwrap()
.unwrap();
assert_eq!(
manager
.async_write_secret("key", &SecretValue::new("value"), None)
.await
.unwrap()
.version_id(),
Some("created")
);
}
#[rstest]
#[case::success(200)]
#[case::denied(403)]
#[tokio::test]
async fn direct_replication_returns_response_or_service_error(#[case] status: u16) {
let server = MockServer::start().await;
scripted_actions(
&server,
vec![Action {
operation: "ReplicateSecretToRegions",
request: json!({"SecretId":"key", "AddReplicaRegions":[{"Region":"region-a"}, {"Region":"region-b"}]}),
status,
response: if status == 200 {
json!({"ARN":"replicated-arn"})
} else {
json!({"__type":"AccessDeniedException"})
},
}],
).await;
let result = manager(&server, Default::default())
.async_replicate_secret("key", &["region-a".into(), "region-b".into()])
.await;
if status == 200 {
assert_eq!(result.unwrap().unwrap().arn(), Some("replicated-arn"));
} else {
assert!(matches!(result, Err(Error::Replicate(_))));
}
}
#[rstest]
#[case::create(false)]
#[case::replicate(true)]
#[tokio::test]
async fn write_and_replication_timeouts_remain_errors(#[case] replicate: bool) {
let server = MockServer::start().await;
Mock::given(wiremock::matchers::method("POST"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(1))
.set_body_json(json!({})),
)
.expect(if replicate { 1 } else { 2 })
.mount(&server)
.await;
let client = Client::from_conf(
client_builder(&server)
.timeout_config(
aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder()
.operation_timeout(Duration::from_millis(50))
.build(),
)
.build(),
);
let manager = AwsSecretsManagerV2::new(client, Default::default());
if replicate {
assert!(matches!(
manager
.async_replicate_secret("key", &["region".into()])
.await,
Err(Error::Replicate(_))
));
} else {
assert!(matches!(
manager
.async_write_secret("key", &SecretValue::new("value"), None)
.await,
Err(Error::Create(_))
));
}
}
#[rstest]
#[case::text("value")]
#[case::json(r#"{"api_key":"test","metadata":{"team":"test"},"temperature":0.7}"#)]
#[case::empty("")]
#[case::unicode(" π\n ")]
#[tokio::test]
async fn write_read_delete_preserves_the_complete_secret_string(#[case] value: &str) {
let server = MockServer::start().await;
scripted_actions(
&server,
vec![
Action {
operation: "CreateSecret",
request: json!({"Name":"key", "SecretString":value, "Description":"description"}),
status: 200,
response: json!({"Name":"key"}),
},
Action {
operation: "GetSecretValue",
request: json!({"SecretId":"key"}),
status: 200,
response: json!({"SecretString":value}),
},
Action {
operation: "DeleteSecret",
request: json!({"SecretId":"key", "RecoveryWindowInDays":7}),
status: 200,
response: json!({"Name":"key"}),
},
],
)
.await;
let manager = manager(&server, Default::default());
assert_eq!(
manager
.async_write_secret("key", &SecretValue::new(value), Some("description"))
.await
.unwrap()
.name(),
Some("key")
);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
value
);
assert_eq!(
manager
.async_delete_secret("key", Some(7))
.await
.unwrap()
.name(),
Some("key")
);
}
#[rstest]
#[case::restore(RecoveryFailure::Restore)]
#[case::update(RecoveryFailure::Update)]
#[case::tag(RecoveryFailure::Tag)]
#[case::delete_after_update(RecoveryFailure::DeleteAfterUpdate)]
#[case::delete_after_tag(RecoveryFailure::DeleteAfterTag)]
#[tokio::test]
async fn failed_update_reschedules_deletion_of_a_restored_alias(#[case] failure: RecoveryFailure) {
let server = MockServer::start().await;
let response = |failed| {
if failed {
(400, json!({"__type":"InvalidRequestException"}))
} else {
(200, json!({}))
}
};
let (restore_status, restore_body) = response(matches!(failure, RecoveryFailure::Restore));
let (update_status, update_body) = response(matches!(
failure,
RecoveryFailure::Update | RecoveryFailure::DeleteAfterUpdate
));
let (delete_status, delete_body) = response(matches!(
failure,
RecoveryFailure::DeleteAfterUpdate | RecoveryFailure::DeleteAfterTag
));
let prefix = [
Action {
operation: "CreateSecret",
request: json!({"Name":"key", "SecretString":"new", "Tags":[{"Key":"stage", "Value":"test"}]}),
status: 400,
response: json!({"__type":"ResourceExistsException"}),
},
Action {
operation: "DescribeSecret",
request: json!({"SecretId":"key"}),
status: 200,
response: json!({"DeletedDate":1}),
},
Action {
operation: "RestoreSecret",
request: json!({"SecretId":"key"}),
status: restore_status,
response: restore_body,
},
];
let update = (!matches!(failure, RecoveryFailure::Restore)).then_some(Action {
operation: "UpdateSecret",
request: json!({"SecretId":"key", "SecretString":"new"}),
status: update_status,
response: update_body,
});
let tag = matches!(
failure,
RecoveryFailure::Tag | RecoveryFailure::DeleteAfterTag
)
.then_some(Action {
operation: "TagResource",
request: json!({"SecretId":"key", "Tags":[{"Key":"stage", "Value":"test"}]}),
status: 400,
response: json!({"__type":"InvalidRequestException"}),
});
let delete = (!matches!(failure, RecoveryFailure::Restore)).then_some(Action {
operation: "DeleteSecret",
request: json!({"SecretId":"key", "RecoveryWindowInDays":7}),
status: delete_status,
response: delete_body,
});
scripted_actions(
&server,
prefix
.into_iter()
.chain(update)
.chain(tag)
.chain(delete)
.collect(),
)
.await;
let error = manager(
&server,
KeyManagementSettings {
tags: Some(std::collections::BTreeMap::from([(
"stage".into(),
"test".into(),
)])),
..Default::default()
},
)
.async_write_secret("key", &SecretValue::new("new"), None)
.await
.unwrap_err();
match failure {
RecoveryFailure::Restore => assert!(matches!(error, Error::Restore(_))),
RecoveryFailure::Update => assert!(matches!(error, Error::Update(_))),
RecoveryFailure::Tag => assert!(matches!(error, Error::Tag(_))),
RecoveryFailure::DeleteAfterUpdate | RecoveryFailure::DeleteAfterTag => {
assert!(matches!(error, Error::Delete(_)))
}
}
}

View file

@ -0,0 +1 @@
- https://learn.microsoft.com/en-us/rest/api/keyvault/secrets/get-secret/get-secret

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
tokio.workspace = true
litellm-auth-azure.workspace = true
litellm-auth-types.workspace = true
litellm-secrets-types.workspace = true
@ -17,7 +18,6 @@ veil.workspace = true
percent-encoding = "2.3"
[dev-dependencies]
tokio.workspace = true
wiremock = "0.6.5"
rstest.workspace = true
serde_json.workspace = true

View file

@ -1,5 +1,9 @@
#[derive(thiserror::Error, veil::Redact)]
pub enum Error {
#[error(transparent)]
Operation(#[from] litellm_secrets_types::Error),
#[error("secret manager operation timed out")]
Timeout,
#[error("{0} environment variable is missing")]
MissingEnvironment(&'static str),
#[error("AZURE_KEY_VAULT_URI is not a valid https vault URL")]

View file

@ -3,7 +3,7 @@ use std::sync::Arc;
use litellm_auth_azure::{AzureAuthInputs, AzureAuthService, ConfigValue};
use litellm_auth_types::{InputSource, Sourced};
use litellm_core_utils::settings::Lookup;
use litellm_secrets_types::{Secret, SecretValue};
use litellm_secrets_types::{AzureOperationContext, BaseSecretManager, Secret, SecretValue};
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC};
use serde::Deserialize;
@ -77,6 +77,12 @@ impl AzureKeyVault {
}
pub async fn get_secret(&self, name: &str) -> Result<Option<Secret>, Error> {
BaseSecretManager::async_read_secret(self, name, &AzureOperationContext::default())
.await
.map(|value| value.map(Secret::String))
}
async fn read(&self, name: &str) -> Result<Option<SecretValue>, Error> {
let token = self
.auth
.get_azure_ad_token(&self.inputs, &|key| self.environment.get(key))
@ -91,6 +97,7 @@ impl AzureKeyVault {
.client
.get(url)
.bearer_auth(token.value().secret().expose())
.header(reqwest::header::ACCEPT, "application/json")
.send()
.await
.map_err(Error::Http)?;
@ -102,7 +109,7 @@ impl AzureKeyVault {
}
let payload: SecretResponse = response.json().await.map_err(Error::Http)?;
let value = payload.value.ok_or(Error::MissingValue)?;
Ok(Some(Secret::String(SecretValue::new(value))))
Ok(Some(SecretValue::new(value)))
}
}
@ -113,3 +120,59 @@ fn scope_for(vault: &reqwest::Url) -> String {
.map_or(host, |(_, remainder)| remainder);
format!("https://{resource}/.default")
}
impl BaseSecretManager for AzureKeyVault {
type Error = Error;
type Context = AzureOperationContext;
async fn async_read_secret(
&self,
name: &str,
context: &Self::Context,
) -> Result<Option<SecretValue>, Error> {
match context.timeout {
Some(timeout) => tokio::time::timeout(timeout, self.read(name))
.await
.map_err(|_| Error::Timeout)?,
None => self.read(name).await,
}
}
}
pub trait AzureTokenProvider: Send + Sync {
fn get_token<'a>(
&'a self,
scope: &'a str,
environment: &'a (dyn Lookup + Send + Sync),
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<SecretValue, Error>> + Send + 'a>>;
}
#[derive(Default)]
pub struct NativeAzureTokenProvider {
auth: AzureAuthService,
}
impl AzureTokenProvider for NativeAzureTokenProvider {
fn get_token<'a>(
&'a self,
scope: &'a str,
environment: &'a (dyn Lookup + Send + Sync),
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<SecretValue, Error>> + Send + 'a>>
{
Box::pin(async move {
let inputs = AzureAuthInputs {
azure_scope: ConfigValue::Value(Sourced::new(
scope.to_owned(),
InputSource::Deployment,
)),
enable_azure_ad_token_refresh: Sourced::new(true, InputSource::Deployment),
..Default::default()
};
self.auth
.get_azure_ad_token(&inputs, &|name| environment.get(name))
.await?
.map(|token| SecretValue::new(token.value().secret().expose()))
.ok_or(Error::MissingCredentials)
})
}
}

View file

@ -4,4 +4,4 @@ mod error;
mod key_vault;
pub use error::Error;
pub use key_vault::AzureKeyVault;
pub use key_vault::{AzureKeyVault, AzureTokenProvider, NativeAzureTokenProvider};

View file

@ -25,6 +25,7 @@ async fn reads_secret_with_bearer_token_and_api_version() {
Mock::given(path("/secrets/OPENAI-API-KEY"))
.and(query_param("api-version", "7.4"))
.and(header("authorization", "Bearer fake"))
.and(header("accept", "application/json"))
.respond_with(
ResponseTemplate::new(200)
.set_body_json(serde_json::json!({"value": "s3cret", "id": "secret-id"})),
@ -42,6 +43,23 @@ async fn reads_secret_with_bearer_token_and_api_version() {
assert_eq!(secret, Secret::String(SecretValue::new("s3cret")));
}
#[rstest]
#[tokio::test]
async fn preserves_secret_contents_and_redacts_debug_output() {
let server = MockServer::start().await;
let value = " \tvalue-π\n";
Mock::given(path("/secrets/NAME"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({"value": value})))
.expect(1)
.mount(&server)
.await;
let secret = manager(&server).get_secret("NAME").await.unwrap().unwrap();
assert_eq!(secret.as_str(), Some(value));
assert!(!format!("{secret:?}").contains(value));
}
#[rstest]
#[tokio::test]
async fn percent_encodes_secret_name_path_segment() {
@ -230,3 +248,22 @@ async fn parity_fixture_matches_python_backend_contract(parity_fixture: Fixture)
}
}
}
#[tokio::test]
async fn trait_read_limits_the_operation_duration() {
use litellm_secrets_types::{AzureOperationContext, BaseSecretManager};
use std::time::Duration;
let server = MockServer::start().await;
Mock::given(wiremock::matchers::method("GET"))
.respond_with(ResponseTemplate::new(200).set_delay(Duration::from_secs(1)))
.mount(&server)
.await;
let manager = manager(&server);
let context = AzureOperationContext {
timeout: Some(Duration::from_millis(30)),
};
assert!(matches!(
BaseSecretManager::async_read_secret(&manager, "key", &context).await,
Err(Error::Timeout)
));
}

View file

@ -0,0 +1 @@
- https://docs.cyberark.com/conjur-open-source/latest/en/content/developer/conjur_api_retrieve_secret.htm

View file

@ -19,7 +19,9 @@ percent-encoding = "2.3"
tokio = { workspace = true, features = ["sync"] }
[dev-dependencies]
rcgen = "0.14.10"
rstest.workspace = true
tempfile = "3.27.0"
tokio.workspace = true
wiremock = "0.6.5"
serde.workspace = true

View file

@ -1,5 +1,7 @@
#[derive(thiserror::Error, veil::Redact)]
pub enum Error {
#[error("CyberArk Conjur operation timed out")]
Timeout,
#[error("CyberArk Conjur HTTP request failed")]
Http(
#[from]

View file

@ -4,4 +4,4 @@ mod error;
mod secret_manager;
pub use error::Error;
pub use secret_manager::{CyberArkSecretManager, DeleteOutcome};
pub use secret_manager::{AuthenticationRetry, CyberArkSecretManager, DeleteOutcome, WriteFailure};

View file

@ -1,9 +1,14 @@
mod client;
mod read;
mod write;
use std::{fs, sync::Arc, time::Duration};
use base64::{Engine, engine::general_purpose::STANDARD};
use litellm_core_utils::settings::Lookup;
use litellm_secrets_types::{
BaseSecretManager, SecretOperationContext, SecretValue, SecretWriteContext,
BaseSecretManager, CyberarkOperationContext, RotationError, SecretCache, SecretDeleter,
SecretRotator, SecretValue, SecretWriteContext, SecretWriter, async_rotate_secret,
validate_secret_name,
};
use moka::future::Cache;
@ -23,6 +28,7 @@ const DEFAULT_API_BASE: &str = "http://127.0.0.1:8080";
const DEFAULT_ACCOUNT: &str = "default";
const DEFAULT_USERNAME: &str = "admin";
const DEFAULT_REFRESH_INTERVAL: Duration = Duration::from_secs(300);
const MAX_TOKEN_LIFETIME: Duration = Duration::from_secs(7 * 60);
const SECRET_NAME_SAFE: &AsciiSet = &NON_ALPHANUMERIC
.remove(b'-')
.remove(b'_')
@ -37,7 +43,7 @@ pub struct CyberArkSecretManager {
username: String,
api_key: SecretValue,
token: Cache<(), SecretValue>,
secrets: Cache<String, SecretValue>,
secrets: SecretCache<String, SecretValue>,
authentication_lock: Arc<tokio::sync::Mutex<()>>,
}
@ -46,346 +52,53 @@ pub enum DeleteOutcome {
NotSupported,
}
impl CyberArkSecretManager {
pub fn with_client(
client: reqwest::Client,
endpoint: reqwest::Url,
account: String,
username: String,
api_key: SecretValue,
refresh_interval: Option<Duration>,
) -> Self {
let endpoint = normalize_endpoint(endpoint);
let ttl = refresh_interval
.filter(|interval| !interval.is_zero())
.unwrap_or(DEFAULT_REFRESH_INTERVAL);
let token = Cache::builder().time_to_live(ttl).build();
let secrets = Cache::builder().time_to_live(ttl).build();
#[derive(Clone, Copy)]
pub enum AuthenticationRetry {
Never,
Unauthorized,
}
#[derive(veil::Redact)]
pub struct WriteFailure {
pub source: Error,
#[redact]
pub request_url: Option<reqwest::Url>,
pub authentication: bool,
}
impl WriteFailure {
fn local(source: Error) -> Self {
Self {
client,
endpoint,
account,
username,
api_key,
token,
secrets,
authentication_lock: Arc::new(tokio::sync::Mutex::new(())),
source,
request_url: None,
authentication: false,
}
}
pub fn new(
environment: Arc<dyn Lookup + Send + Sync>,
enterprise_enabled: bool,
) -> Result<Self, Error> {
let api_key = environment.get(CYBERARK_API_KEY).unwrap_or_default();
let cert = environment.get(CYBERARK_CLIENT_CERT).unwrap_or_default();
let key = environment.get(CYBERARK_CLIENT_KEY).unwrap_or_default();
if api_key.is_empty() && (cert.is_empty() || key.is_empty()) {
return Err(Error::MissingCredentials);
fn request(source: Error, url: reqwest::Url) -> Self {
Self {
source,
request_url: Some(url),
authentication: false,
}
if !enterprise_enabled {
return Err(Error::EnterpriseRequired);
}
let verify = environment
.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 endpoint = reqwest::Url::parse(
&environment
.get(CYBERARK_API_BASE)
.unwrap_or_else(|| DEFAULT_API_BASE.to_owned()),
)
.map_err(|_| Error::Endpoint)?;
let account = environment
.get(CYBERARK_ACCOUNT)
.unwrap_or_else(|| DEFAULT_ACCOUNT.to_owned());
let username = environment
.get(CYBERARK_USERNAME)
.unwrap_or_else(|| DEFAULT_USERNAME.to_owned());
let refresh_interval = environment
.get(CYBERARK_REFRESH_INTERVAL)
.map(|value| {
value
.parse::<u64>()
.map(Duration::from_secs)
.map_err(|_| Error::RefreshInterval)
})
.transpose()?;
Ok(Self::with_client(
client,
endpoint,
account,
username,
SecretValue::new(api_key),
refresh_interval,
))
}
}
impl CyberArkSecretManager {
fn secret_url(&self, name: &str) -> Result<reqwest::Url, Error> {
let encoded = utf8_percent_encode(name, SECRET_NAME_SAFE);
self.endpoint
.join(&format!("secrets/{}/variable/{}", self.account, encoded))
.map_err(|_| Error::Endpoint)
}
async fn authenticate(&self, context: &SecretOperationContext) -> Result<SecretValue, Error> {
if let Some(token) = self.token.get(&()).await {
return Ok(token);
}
let _guard = self.authentication_lock.lock().await;
if let Some(token) = self.token.get(&()).await {
return Ok(token);
}
let url = self
.endpoint
.join(&format!(
"authn/{}/{}/authenticate",
self.account, self.username
))
.map_err(|_| Error::Endpoint)?;
let response = with_timeout(
self.client.post(url).body(self.api_key.expose().to_owned()),
context,
)
.send()
.await?;
if !response.status().is_success() {
return Err(Error::AuthStatus(response.status().as_u16()));
}
let token = SecretValue::new(STANDARD.encode(response.text().await?));
self.token.insert((), token.clone()).await;
Ok(token)
}
async fn authorization_header(
&self,
context: &SecretOperationContext,
) -> Result<String, Error> {
Ok(format!(
"Token token=\"{}\"",
self.authenticate(context).await?.expose()
))
}
pub async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
self.async_read_secret_with_context(name, &SecretOperationContext::default())
.await
}
pub async fn async_read_secret_with_context(
&self,
name: &str,
context: &SecretOperationContext,
) -> Result<Option<SecretValue>, Error> {
if let Some(value) = self.secrets.get(name).await {
return Ok(Some(value));
}
let response = with_timeout(
self.client
.get(self.secret_url(name)?)
.header("Authorization", self.authorization_header(context).await?),
context,
)
.send()
.await?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Ok(None);
}
if !response.status().is_success() {
return Err(Error::Status(response.status().as_u16()));
}
let value = SecretValue::new(response.text().await?);
self.secrets.insert(name.to_owned(), value.clone()).await;
Ok(Some(value))
}
pub async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
description: Option<&str>,
) -> Result<(), Error> {
self.async_write_secret_with_context(
name,
value,
description,
&SecretOperationContext::default(),
)
.await
}
pub async fn async_write_secret_with_context(
&self,
name: &str,
value: &SecretValue,
_description: Option<&str>,
context: &SecretOperationContext,
) -> Result<(), Error> {
validate_secret_name(name)?;
self.ensure_variable_exists(name, context).await;
let response = with_timeout(
self.client
.post(self.secret_url(name)?)
.header("Authorization", self.authorization_header(context).await?)
.body(value.expose().to_owned()),
context,
)
.send()
.await?;
if !response.status().is_success() {
return Err(Error::Status(response.status().as_u16()));
}
self.secrets.insert(name.to_owned(), value.clone()).await;
Ok(())
}
async fn ensure_variable_exists(&self, name: &str, context: &SecretOperationContext) {
let policy_url = self
.endpoint
.join(&format!("policies/{}/policy/root", self.account));
let Ok(policy_url) = policy_url else {
litellm_tracing::warn!("Could not build CyberArk policy endpoint");
return;
};
let Ok(authorization) = self.authorization_header(context).await else {
litellm_tracing::warn!(
"Could not authenticate while ensuring CyberArk variable exists"
);
return;
};
let body = format!(
"- !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}");
}
}
}
pub async fn async_delete_secret(
&self,
name: &str,
recovery_window_in_days: Option<u32>,
) -> Result<DeleteOutcome, Error> {
self.async_delete_secret_with_context(
name,
recovery_window_in_days,
&SecretOperationContext::default(),
)
.await
}
pub async fn async_delete_secret_with_context(
&self,
name: &str,
_recovery_window_in_days: Option<u32>,
_context: &SecretOperationContext,
) -> Result<DeleteOutcome, Error> {
litellm_tracing::warn!(
"CyberArk Conjur does not support direct secret deletion. Secrets must be removed through policy updates."
);
self.secrets.invalidate(name).await;
Ok(DeleteOutcome::NotSupported)
}
}
impl BaseSecretManager for CyberArkSecretManager {
type Error = Error;
type WriteResponse = ();
type DeleteResponse = DeleteOutcome;
async fn async_read_secret(
&self,
name: &str,
context: &SecretOperationContext,
) -> Result<Option<SecretValue>, Error> {
self.async_read_secret_with_context(name, context).await
}
async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
context: &SecretWriteContext,
) -> Result<(), Error> {
self.async_write_secret_with_context(
name,
value,
context.description.as_deref(),
&context.operation,
)
.await
}
async fn async_delete_secret(
&self,
name: &str,
recovery_window_in_days: Option<u32>,
context: &SecretOperationContext,
) -> Result<DeleteOutcome, Error> {
self.async_delete_secret_with_context(name, recovery_window_in_days, context)
.await
}
}
fn with_timeout(
request: reqwest::RequestBuilder,
context: &SecretOperationContext,
context: &CyberarkOperationContext,
) -> reqwest::RequestBuilder {
match context.timeout() {
match context.timeout {
Some(timeout) => request.timeout(timeout),
None => request,
}
}
fn normalize_endpoint(mut endpoint: reqwest::Url) -> reqwest::Url {
if !endpoint.path().ends_with('/') {
endpoint.set_path(&format!("{}/", endpoint.path()));
}
endpoint
}

View file

@ -0,0 +1,147 @@
use super::*;
impl CyberArkSecretManager {
pub fn with_client(
client: reqwest::Client,
endpoint: reqwest::Url,
account: String,
username: String,
api_key: SecretValue,
refresh_interval: Option<Duration>,
) -> Self {
let endpoint = normalize_endpoint(endpoint);
let ttl = refresh_interval
.filter(|interval| !interval.is_zero())
.unwrap_or(DEFAULT_REFRESH_INTERVAL);
let token = Cache::builder()
.time_to_live(ttl.min(MAX_TOKEN_LIFETIME))
.build();
let secrets = SecretCache::new(200, ttl);
Self {
client,
endpoint,
account,
username,
api_key,
token,
secrets,
authentication_lock: Arc::new(tokio::sync::Mutex::new(())),
}
}
pub fn new(
environment: Arc<dyn Lookup + Send + Sync>,
enterprise_enabled: bool,
) -> Result<Self, Error> {
let api_key = environment.get(CYBERARK_API_KEY).unwrap_or_default();
let cert = environment.get(CYBERARK_CLIENT_CERT).unwrap_or_default();
let key = environment.get(CYBERARK_CLIENT_KEY).unwrap_or_default();
if api_key.is_empty() && (cert.is_empty() || key.is_empty()) {
return Err(Error::MissingCredentials);
}
if !enterprise_enabled {
return Err(Error::EnterpriseRequired);
}
let verify = environment
.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 endpoint = reqwest::Url::parse(
&environment
.get(CYBERARK_API_BASE)
.unwrap_or_else(|| DEFAULT_API_BASE.to_owned()),
)
.map_err(|_| Error::Endpoint)?;
let account = environment
.get(CYBERARK_ACCOUNT)
.unwrap_or_else(|| DEFAULT_ACCOUNT.to_owned());
let username = environment
.get(CYBERARK_USERNAME)
.unwrap_or_else(|| DEFAULT_USERNAME.to_owned());
let refresh_interval = environment
.get(CYBERARK_REFRESH_INTERVAL)
.map(|value| {
value
.parse::<u64>()
.map(Duration::from_secs)
.map_err(|_| Error::RefreshInterval)
})
.transpose()?;
Ok(Self::with_client(
client,
endpoint,
account,
username,
SecretValue::new(api_key),
refresh_interval,
))
}
pub(super) fn authentication_url(&self) -> Result<reqwest::Url, Error> {
self.endpoint
.join(&format!(
"authn/{}/{}/authenticate",
self.account,
utf8_percent_encode(&self.username, SECRET_NAME_SAFE)
))
.map_err(|_| Error::Endpoint)
}
pub(super) async fn authenticate(
&self,
context: &CyberarkOperationContext,
) -> Result<SecretValue, Error> {
if let Some(token) = self.token.get(&()).await {
return Ok(token);
}
let _guard = self.authentication_lock.lock().await;
if let Some(token) = self.token.get(&()).await {
return Ok(token);
}
let url = self.authentication_url()?;
let response = with_timeout(
self.client.post(url).body(self.api_key.expose().to_owned()),
context,
)
.send()
.await?;
if !response.status().is_success() {
return Err(Error::AuthStatus(response.status().as_u16()));
}
let token = SecretValue::new(STANDARD.encode(response.text().await?));
self.token.insert((), token.clone()).await;
Ok(token)
}
pub(super) async fn authorization_header(
&self,
context: &CyberarkOperationContext,
) -> Result<String, Error> {
Ok(format!(
"Token token=\"{}\"",
self.authenticate(context).await?.expose()
))
}
}
fn normalize_endpoint(mut endpoint: reqwest::Url) -> reqwest::Url {
if !endpoint.path().ends_with('/') {
endpoint.set_path(&format!("{}/", endpoint.path()));
}
endpoint
}

View file

@ -0,0 +1,118 @@
use super::*;
impl CyberArkSecretManager {
pub async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
self.async_read_secret_with_context(name, &CyberarkOperationContext::default())
.await
}
pub async fn async_read_secret_with_context(
&self,
name: &str,
context: &CyberarkOperationContext,
) -> Result<Option<SecretValue>, Error> {
self.read_with_retry(name, context, AuthenticationRetry::Unauthorized)
.await
}
pub async fn read_with_retry(
&self,
name: &str,
context: &CyberarkOperationContext,
retry: AuthenticationRetry,
) -> Result<Option<SecretValue>, Error> {
validate_secret_name(name)?;
let read = self.secrets.read(
name.to_owned(),
self.read_uncached_with_retry(name, context, retry),
);
match context.timeout {
Some(timeout) => tokio::time::timeout(timeout, read)
.await
.map_err(|_| Error::Timeout)?,
None => read.await,
}
}
pub(super) async fn read_uncached(
&self,
name: &str,
context: &CyberarkOperationContext,
) -> Result<Option<SecretValue>, Error> {
self.read_uncached_with_retry(name, context, AuthenticationRetry::Unauthorized)
.await
}
pub async fn read_fresh_with_retry(
&self,
name: &str,
context: &CyberarkOperationContext,
retry: AuthenticationRetry,
) -> Result<Option<SecretValue>, Error> {
validate_secret_name(name)?;
self.secrets
.refresh(
name.to_owned(),
self.read_uncached_with_retry(name, context, retry),
)
.await
}
pub async fn invalidate_cached_secret(&self, name: &str) {
self.secrets.invalidate(&name.to_owned()).await;
}
pub(super) async fn read_uncached_with_retry(
&self,
name: &str,
context: &CyberarkOperationContext,
retry: AuthenticationRetry,
) -> Result<Option<SecretValue>, Error> {
let had_cached_token = self.token.get(&()).await.is_some();
let response = with_timeout(
self.client
.get(self.secret_url(name)?)
.header("Authorization", self.authorization_header(context).await?),
context,
)
.send()
.await?;
let response = if matches!(retry, AuthenticationRetry::Unauthorized)
&& had_cached_token
&& response.status() == reqwest::StatusCode::UNAUTHORIZED
{
self.token.invalidate(&()).await;
with_timeout(
self.client
.get(self.secret_url(name)?)
.header("Authorization", self.authorization_header(context).await?),
context,
)
.send()
.await?
} else {
response
};
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Ok(None);
}
if !response.status().is_success() {
return Err(Error::Status(response.status().as_u16()));
}
let value = SecretValue::new(response.text().await?);
Ok(Some(value))
}
}
impl BaseSecretManager for CyberArkSecretManager {
type Error = Error;
type Context = CyberarkOperationContext;
async fn async_read_secret(
&self,
name: &str,
context: &Self::Context,
) -> Result<Option<SecretValue>, Error> {
self.async_read_secret_with_context(name, context).await
}
}

Some files were not shown because too many files have changed in this diff Show more