Merge remote-tracking branch 'origin/main' into litellm_langfuse_sdk_v4
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Modules / fmt, validate, test (aws) (push) Has been cancelled
Terraform Modules / fmt, validate, test (gcp) (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

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

# Conflicts:
#	tests/integration/contracts.json
This commit is contained in:
yucheng 2026-09-23 20:39:08 +00:00
commit 09cc332c03
346 changed files with 32327 additions and 7027 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 >>
@ -3231,7 +3231,7 @@ workflows:
name: integration-<< matrix.suite >>
matrix:
parameters:
suite: [management, accounting, database, providers, extensions, sdk, cost, browser]
suite: [management, accounting, database, providers, extensions, mcp, sdk, cost, browser]
filters:
branches:
only:

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="main"
merge_base=""
for base in $candidate_bases; do
git fetch --quiet origin "$base" 2>/dev/null || continue

View file

@ -111,6 +111,15 @@ upstream_pid=$!
if [ "$suite" = cost ]; then
export INTEGRATION_WORKERS=8
fi
if [ "$suite" = mcp ]; then
export INTEGRATION_WORKERS=4 INTEGRATION_COVERAGE=1
fi
coverage_data="$PWD/$results/coverage/data"
proxy_command=(.venv/bin/python -m integration._support.proxy)
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then
mkdir -p "$(dirname "$coverage_data")"
proxy_command=(.venv/bin/python -m coverage run --rcfile=tests/integration/mcp_coverage.toml -m integration._support.proxy)
fi
start_proxy() {
local port="$1"
local log_name="$2"
@ -133,8 +142,8 @@ start_proxy() {
INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \
LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" LITELLM_UI_PATH="$LITELLM_UI_PATH" PROXY_BASE_URL="http://127.0.0.1:$port" \
LITELLM_MODE=PRODUCTION STORE_MODEL_IN_DB=True "${cost_map_env[@]}" \
AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
.venv/bin/python -m integration._support.proxy --config tests/integration/proxy_config.yaml \
AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 COVERAGE_FILE="$coverage_data" \
"${proxy_command[@]}" --config tests/integration/proxy_config.yaml \
--host 127.0.0.1 --port "$port" --num_workers 1 --telemetry False \
--use_prisma_db_push --enforce_prisma_migration_check \
> "$results/$log_name" 2>&1 &
@ -146,7 +155,7 @@ proxy_pid="$launched_pid"
curl --noproxy '*' -sSf -X POST "$INTEGRATION_PROXY_URL/config/update" \
-H "Authorization: Bearer $LITELLM_MASTER_KEY" -H 'Content-Type: application/json' \
-d '{"router_settings": {"num_retries": 0}}' > "$results/seed-router-settings.json"
if [ "$suite" = management ]; then
if [ "$suite" = management ] || [ "$suite" = mcp ]; then
export INTEGRATION_PEER_URL=http://127.0.0.1:4001
start_proxy 4001 peer.log
peer_pid="$launched_pid"
@ -187,3 +196,23 @@ env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
INTEGRATION_ORDER_SEED="$INTEGRATION_ORDER_SEED" \
LITELLM_LOCAL_MODEL_COST_MAP=True AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
.venv/bin/python tests/integration/run.py "$suite" --results "$results"
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then
for covered_pid in "$proxy_pid" "$peer_pid"; do
[ -n "$covered_pid" ] || continue
kill -TERM -- "-$covered_pid"
for _ in {1..300}; do
kill -0 "$covered_pid" 2>/dev/null || break
sleep 0.1
done
wait "$covered_pid" 2>/dev/null || true
done
proxy_pid=""
peer_pid=""
COVERAGE_FILE="$coverage_data" .venv/bin/python -m coverage combine --rcfile=tests/integration/mcp_coverage.toml
COVERAGE_FILE="$coverage_data" .venv/bin/python -m coverage report --rcfile=tests/integration/mcp_coverage.toml \
> "$results/coverage/coverage.txt"
COVERAGE_FILE="$coverage_data" .venv/bin/python -m coverage html --rcfile=tests/integration/mcp_coverage.toml \
-d "$results/coverage/html"
tail -n 1 "$results/coverage/coverage.txt"
fi

View file

@ -31,8 +31,8 @@ def main() -> None:
result: Final = json.loads(Path(sys.argv[1]).read_text())
assert not result.get("errors"), result.get("errors")
expected: Final = json.loads(
(Path(__file__).resolve().parents[2] / "tests/integration/contracts.json").read_text()
)["browser"]
(Path(__file__).resolve().parents[2] / "tests/e2e/ui/tests/integrationCritical/expected.json").read_text()
)
assert expected and result["stats"]["expected"] == len(expected)
assert all(result["stats"][name] == 0 for name in ("unexpected", "flaky", "skipped"))

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/

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

@ -235,9 +235,7 @@ class Slice:
return True # a `-k` this parser cannot model is assumed to claim everything
if any(term.lower() in relative_path.lower() for term in self.excluded):
return False
return not self.required or any(
term.lower() in name.lower() for term in self.required for name in inner_names
)
return not self.required or any(term.lower() in name.lower() for term in self.required for name in inner_names)
def _strings(node: object) -> Iterable[str]:
@ -307,9 +305,7 @@ def _matchable_names(relative_path: str) -> frozenset[str]:
except (OSError, SyntaxError):
return frozenset({relative_path})
return frozenset({relative_path}) | frozenset(
node.name
for node in ast.walk(tree)
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef))
node.name for node in ast.walk(tree) if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef))
)
@ -331,9 +327,7 @@ def _deselected_everywhere(allowlist: Allowlist) -> tuple[Finding, ...]:
slices: Final = _slices()
named_by_workflow: Final = _workflow_named_tokens()
globbed: Final = tuple(
path
for path in _test_files()
if any(_token_covers(glob, path) for slice_ in slices for glob in slice_.globs)
path for path in _test_files() if any(_token_covers(glob, path) for slice_ in slices for glob in slice_.globs)
)
return tuple(
Finding(
@ -363,11 +357,7 @@ def _shard_children(root: str, repo_root: pathlib.Path = REPO_ROOT) -> tuple[str
child.relative_to(repo_root).as_posix()
for child in (repo_root / root).iterdir()
if not child.name.startswith(".")
and (
_holds_tests(child)
if child.is_dir()
else child.name.startswith("test_") and child.suffix == ".py"
)
and (_holds_tests(child) if child.is_dir() else child.name.startswith("test_") and child.suffix == ".py")
)
)
@ -499,13 +489,32 @@ def _check_shards() -> int:
return 0
def _integration_groups(runner: pathlib.Path) -> dict[str, tuple[str, ...]]:
module: Final = ast.parse(runner.read_text())
literal: Final = next(
node.value
for node in module.body
if isinstance(node, ast.AnnAssign) and isinstance(node.target, ast.Name) and node.target.id == "GROUPS"
)
mapping: Final = literal.args[0] if isinstance(literal, ast.Call) else literal
return {group: tuple(folders) for group, folders in ast.literal_eval(mapping).items()}
def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozenset[str], tuple[Finding, ...]]:
manifest: Final = repo_root / "tests/integration/contracts.json"
if not manifest.exists():
runner: Final = repo_root / "tests/integration/run.py"
if not runner.exists():
return frozenset(), ()
entries: Final = json.loads(manifest.read_text())
paths: Final = frozenset(node.split("::", 1)[0] for node in entries["tests"])
browser_paths: Final = frozenset(node.split("::", 1)[0] for node in entries.get("browser", {}))
groups: Final = _integration_groups(runner)
integration_root: Final = repo_root / "tests/integration"
paths: Final = frozenset(
str(path.relative_to(repo_root))
for folders in groups.values()
for folder in folders
for path in (integration_root / folder).glob("test_*.py")
)
browser_manifest: Final = repo_root / "tests/e2e/ui/tests/integrationCritical/expected.json"
browser_nodes: Final = json.loads(browser_manifest.read_text()) if browser_manifest.exists() else ()
browser_paths: Final = frozenset(node.split("::", 1)[0] for node in browser_nodes)
circle_path: Final = repo_root / ".circleci/config.yml"
circle: Final = yaml.safe_load(circle_path.read_text()) if circle_path.exists() else {}
steps: Final = circle.get("jobs", {}).get("integration_contracts", {}).get("steps", ())
@ -526,15 +535,14 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
)
required: Final = (frozenset({"browser"}) if browser_paths else frozenset()) | frozenset(
group
for group, folders in entries["groups"].items()
for group, folders in groups.items()
if any(any(path.startswith(f"tests/integration/{folder}/") for folder in folders) for path in paths)
)
ungrouped: Final = frozenset(
path
for path in paths
if sum(
any(path.startswith(f"tests/integration/{folder}/") for folder in folders)
for folders in entries["groups"].values()
any(path.startswith(f"tests/integration/{folder}/") for folder in folders) for folders in groups.values()
)
!= 1
)
@ -547,10 +555,6 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
Finding(path, "integration contract is also selected by GitHub Actions")
for path in paths
if any(_token_covers(token, path) for token in gha_tokens)
) + tuple(
Finding(path, "canonical integration test file is missing")
for path in paths
if not (repo_root / path).is_file()
)
browser_commands: Final = tuple(
scalar.value
@ -592,7 +596,7 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
) + tuple(Finding(path, "canonical node must have exactly one integration group") for path in sorted(ungrouped))
if not paths or not invoked or not scheduled:
return frozenset(), findings + (
Finding(str(manifest.relative_to(repo_root)), "dedicated CircleCI runner is missing"),
Finding(str(runner.relative_to(repo_root)), "dedicated CircleCI runner is missing"),
)
return paths | browser_paths, findings + group_findings + browser_findings + exclusion_findings

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,186 +0,0 @@
name: Create Release
on:
workflow_dispatch:
inputs:
tag:
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0-dev.2, 1.84.0.post1; legacy v1.83.10-stable still accepted)"
required: true
type: string
commit_hash:
description: "Full 40-char commit SHA to target"
required: true
type: string
permissions: {}
jobs:
release:
name: Create Release
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- name: Validate inputs
env:
TAG: ${{ inputs.tag }}
COMMIT_HASH: ${{ inputs.commit_hash }}
run: |
if ! echo "${COMMIT_HASH}" | grep -qE '^[0-9a-f]{40}$'; then
echo "::error::commit_hash must be a full 40-character commit SHA"
exit 1
fi
if ! echo "${TAG}" | grep -qE '^v?[0-9]+\.[0-9]+\.[0-9]+'; then
echo "::error::tag must start with X.Y.Z (optional leading v), e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, or v1.83.10-stable"
exit 1
fi
- name: Create release
env:
TAG: ${{ inputs.tag }}
COMMIT_HASH: ${{ inputs.commit_hash }}
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
with:
script: |
const tag = process.env.TAG;
const commitHash = process.env.COMMIT_HASH;
// Mark RC / dev / nightly / alpha / beta tags as GitHub pre-releases.
// Accept both PEP 440 (`.dev`) and SemVer (`-dev`) separators so tags
// like `1.84.0.dev2` and `1.84.0-dev.2` are both detected.
// PEP 440 post-releases (e.g. `1.84.0.post1`) and legacy `-stable[.patch.N]`
// are stable maintenance releases, not pre-releases.
const isPrerelease = /(?:rc|nightly|alpha|beta|[-.]dev)/i.test(tag);
// A stable release should only claim the repo "latest" badge when its
// version is >= the current latest. Otherwise a backport (e.g. 1.84.6)
// would steal "latest" from a newer line (e.g. 1.88.1).
const versionKey = (rawTag) => {
const m = String(rawTag).match(/^v?(\d+)\.(\d+)\.(\d+)/);
if (!m) return null;
const maintenance = String(rawTag).match(/(?:\.post|\.patch\.)(\d+)/i);
return [Number(m[1]), Number(m[2]), Number(m[3]), maintenance ? Number(maintenance[1]) : 0];
};
const isAtLeast = (a, b) => {
for (let i = 0; i < a.length; i++) {
if (a[i] !== b[i]) return a[i] > b[i];
}
return true;
};
const cosignSection = [
`## Verify Docker Image Signature`,
``,
`All LiteLLM Docker images are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). Every release is signed with the same key introduced in [commit \`0112e53\`](https://github.com/BerriAI/litellm/commit/0112e53046018d726492c814b3644b7d376029d0).`,
``,
`**Verify using the pinned commit hash (recommended):**`,
``,
`A commit hash is cryptographically immutable, so this is the strongest way to ensure you are using the original signing key:`,
``,
'```bash',
`cosign verify \\`,
` --key https://raw.githubusercontent.com/BerriAI/litellm/0112e53046018d726492c814b3644b7d376029d0/cosign.pub \\`,
` ghcr.io/berriai/litellm:${tag}`,
'```',
``,
`**Verify using the release tag (convenience):**`,
``,
`Tags are protected in this repository and resolve to the same key. This option is easier to read but relies on tag protection rules:`,
``,
'```bash',
`cosign verify \\`,
` --key https://raw.githubusercontent.com/BerriAI/litellm/${tag}/cosign.pub \\`,
` ghcr.io/berriai/litellm:${tag}`,
'```',
``,
`Expected output:`,
``,
'```',
`The following checks were performed on each of these signatures:`,
` - The cosign claims were validated`,
` - The signatures were verified against the specified public key`,
'```',
``,
`---`,
``,
].join('\n');
try {
let makeLatest = "false";
const newVersion = versionKey(tag);
if (!isPrerelease && newVersion) {
let latestVersion = null;
try {
const latest = await github.rest.repos.getLatestRelease({
owner: context.repo.owner,
repo: context.repo.repo,
});
latestVersion = versionKey(latest.data.tag_name);
} catch (error) {
if (error.status !== 404) throw error;
}
makeLatest = (!latestVersion || isAtLeast(newVersion, latestVersion)) ? "true" : "false";
}
try {
await github.rest.git.createRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref: `refs/tags/${tag}`,
sha: commitHash,
});
} catch (error) {
if (error.status !== 422) throw error;
const existing = await github.rest.git.getRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref: `tags/${tag}`,
});
if (existing.data.object.sha !== commitHash) {
throw new Error(`Tag ${tag} already exists at ${existing.data.object.sha}, expected ${commitHash}`);
}
}
const response = await github.rest.repos.createRelease({
draft: true,
generate_release_notes: true,
name: tag,
owner: context.repo.owner,
prerelease: isPrerelease,
repo: context.repo.repo,
tag_name: tag,
});
const updatedBody = cosignSection + (response.data.body ?? '');
await github.rest.repos.updateRelease({
owner: context.repo.owner,
repo: context.repo.repo,
release_id: response.data.id,
tag_name: tag,
body: updatedBody,
draft: false,
});
if (!isPrerelease) {
await github.rest.repos.updateRelease({
owner: context.repo.owner,
repo: context.repo.repo,
release_id: response.data.id,
tag_name: tag,
make_latest: makeLatest,
});
}
} catch (error) {
core.setFailed(error.message);
}
create-branch:
name: Create Release Branch
needs: release
permissions:
contents: write
uses: ./.github/workflows/create-release-branch.yml
with:
tag: ${{ inputs.tag }}
commit_hash: ${{ inputs.commit_hash }}

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,8 +4,6 @@ on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "litellm/_redis.py"

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:

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

@ -26,6 +26,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
"/v2/login",
"/v3/login",
"/logout",
"/session/logout",
"/token",
"/onboarding/",
"/audit",

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

@ -8,10 +8,6 @@ mod logger;
mod marshal;
mod python_settings;
mod routes;
#[allow(
dead_code,
reason = "secret-manager foundations await rollout activation"
)]
mod secrets;
mod tokenizer;
@ -56,7 +52,11 @@ mod _native {
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("_ResponseCacheRuntime", py.get_type::<ResolvedCache>())?;
dict.set_item(
"_SecretManagerRuntime",
py.get_type::<crate::secrets::runtime::NativeSecretManager>(),
)
}
}

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

@ -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
}
}

View file

@ -0,0 +1,262 @@
use super::*;
impl CyberArkSecretManager {
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,
&CyberarkOperationContext::default(),
)
.await
}
pub async fn async_write_secret_with_context(
&self,
name: &str,
value: &SecretValue,
_description: Option<&str>,
context: &CyberarkOperationContext,
) -> Result<(), Error> {
self.write_with_retry(name, value, context, AuthenticationRetry::Unauthorized)
.await
.map_err(|failure| failure.source)
}
pub async fn write_with_retry(
&self,
name: &str,
value: &SecretValue,
context: &CyberarkOperationContext,
retry: AuthenticationRetry,
) -> Result<(), WriteFailure> {
validate_secret_name(name).map_err(|source| WriteFailure::local(source.into()))?;
self.ensure_variable_exists(name, context).await;
let url = self.secret_url(name).map_err(WriteFailure::local)?;
let response = self.post_value(&url, value, context).await?;
let response = if matches!(retry, AuthenticationRetry::Unauthorized)
&& response.status() == reqwest::StatusCode::UNAUTHORIZED
{
self.token.invalidate(&()).await;
self.post_value(&url, value, context).await?
} else {
response
};
if !response.status().is_success() {
return Err(WriteFailure::request(
Error::Status(response.status().as_u16()),
url,
));
}
self.secrets.insert(name.to_owned(), value.clone()).await;
Ok(())
}
async fn post_value(
&self,
url: &reqwest::Url,
value: &SecretValue,
context: &CyberarkOperationContext,
) -> Result<reqwest::Response, WriteFailure> {
let authorization =
self.authorization_header(context)
.await
.map_err(|source| WriteFailure {
source,
request_url: self.authentication_url().ok(),
authentication: true,
})?;
with_timeout(
self.client
.post(url.clone())
.header("Authorization", authorization)
.body(value.expose().to_owned()),
context,
)
.send()
.await
.map_err(|source| WriteFailure::request(source.into(), url.clone()))
}
pub(super) async fn ensure_variable_exists(
&self,
name: &str,
context: &CyberarkOperationContext,
) {
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_rotate_secret(
&self,
current_name: &str,
new_name: &str,
value: &SecretValue,
) -> Result<(), RotationError<(), Error>> {
self.async_rotate_secret_with_context(
current_name,
new_name,
value,
&CyberarkOperationContext::default(),
)
.await
}
pub async fn async_rotate_secret_with_context(
&self,
current_name: &str,
new_name: &str,
value: &SecretValue,
context: &CyberarkOperationContext,
) -> Result<(), RotationError<(), Error>> {
async_rotate_secret(self, current_name, new_name, value, context).await
}
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,
&CyberarkOperationContext::default(),
)
.await
}
pub async fn async_delete_secret_with_context(
&self,
name: &str,
_recovery_window_in_days: Option<u32>,
_context: &CyberarkOperationContext,
) -> 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.to_owned()).await;
Ok(DeleteOutcome::NotSupported)
}
}
impl SecretWriter for CyberArkSecretManager {
type WriteResponse = ();
async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
context: &SecretWriteContext<Self::Context>,
) -> Result<(), Error> {
self.async_write_secret_with_context(
name,
value,
context.description.as_deref(),
&context.operation,
)
.await
}
}
impl SecretDeleter for CyberArkSecretManager {
type DeleteResponse = DeleteOutcome;
async fn async_delete_secret(
&self,
name: &str,
context: &Self::Context,
) -> Result<DeleteOutcome, Error> {
self.async_delete_secret_with_context(name, None, context)
.await
}
}
impl SecretRotator for CyberArkSecretManager {
type RotationResponse = ();
async fn async_read_secret_fresh(
&self,
name: &str,
context: &Self::Context,
) -> Result<Option<SecretValue>, Error> {
validate_secret_name(name)?;
let read = self
.secrets
.refresh(name.to_owned(), self.read_uncached(name, context));
match context.timeout {
Some(timeout) => tokio::time::timeout(timeout, read)
.await
.map_err(|_| Error::Timeout)?,
None => read.await,
}
}
async fn async_write_replacement(
&self,
current_name: &str,
new_name: &str,
value: &SecretValue,
context: &Self::Context,
) -> Result<(), Error> {
SecretWriter::async_write_secret(
self,
new_name,
value,
&SecretWriteContext::rotated_from(current_name, context.clone()),
)
.await
}
}

View file

@ -1,10 +1,14 @@
use std::{sync::Arc, time::Duration};
use std::{
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
time::Duration,
};
use base64::{Engine, engine::general_purpose::STANDARD};
use litellm_secrets_cyberark::{CyberArkSecretManager, DeleteOutcome, Error};
use litellm_secrets_types::{
BaseSecretManager, CyberarkOperationContext, SecretOperationContext, SecretValue,
};
use litellm_secrets_types::{BaseSecretManager, CyberarkOperationContext, SecretValue};
use rstest::{fixture, rstest};
use serde::Deserialize;
use wiremock::{
@ -12,546 +16,13 @@ use wiremock::{
matchers::{body_string, header, method, path},
};
const TOKEN_JSON: &str = r#"{"protected":"p","payload":"q","signature":"s"}"#;
#[path = "secret_manager/support.rs"]
mod support;
use support::*;
#[derive(Deserialize)]
struct ParityFixture {
endpoint: String,
account: String,
username: String,
api_key: String,
authenticate_path: String,
token_json: String,
authorization_header: String,
policy_path: String,
secrets: Vec<ParitySecret>,
}
#[derive(Deserialize)]
struct ParitySecret {
name: String,
path: String,
policy_body: String,
}
#[derive(Debug)]
struct RawPath(String);
impl Match for RawPath {
fn matches(&self, request: &Request) -> bool {
request.url.path() == self.0
}
}
#[fixture]
fn parity_fixture() -> ParityFixture {
serde_json::from_str(include_str!("fixtures/parity.json")).unwrap()
}
fn manager(server: &MockServer, ttl: Duration) -> CyberArkSecretManager {
CyberArkSecretManager::with_client(
reqwest::Client::new(),
server.uri().parse().unwrap(),
"acct".into(),
"admin".into(),
SecretValue::new("k3y"),
Some(ttl),
)
}
async fn mount_auth(server: &MockServer, expected: u64) {
Mock::given(method("POST"))
.and(path("/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.expect(expected)
.mount(server)
.await;
}
#[rstest]
#[tokio::test]
async fn successful_reads_cache_auth_secret_and_redact_values() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
let token = STANDARD.encode(TOKEN_JSON);
Mock::given(path("/secrets/acct/variable/OPENAI_API_KEY"))
.and(header("authorization", format!("Token token=\"{token}\"")))
.respond_with(ResponseTemplate::new(200).set_body_string("sk-live"))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
for _ in 0..2 {
let value = manager
.async_read_secret("OPENAI_API_KEY")
.await
.unwrap()
.unwrap();
assert_eq!(value.expose(), "sk-live");
assert!(!format!("{value:?}").contains("sk-live"));
}
}
#[rstest]
#[tokio::test]
async fn concurrent_reads_share_authentication_request() {
let server = MockServer::start().await;
Mock::given(path("/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(
ResponseTemplate::new(200)
.set_body_string(TOKEN_JSON)
.set_delay(Duration::from_millis(20)),
)
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.and(header(
"authorization",
format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON)),
))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
let (first, second) = tokio::join!(
manager.async_read_secret("key"),
manager.async_read_secret("key")
);
assert_eq!(first.unwrap().unwrap().expose(), "value");
assert_eq!(second.unwrap().unwrap().expose(), "value");
}
#[rstest]
#[case::not_found(404)]
#[case::unauthorized(401)]
#[case::forbidden(403)]
#[case::server_error(500)]
#[tokio::test]
async fn failed_reads_are_not_cached(#[case] status: u16) {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
let failing = Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(status))
.expect(1)
.mount_as_scoped(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
let result = manager.async_read_secret("key").await;
if status == 404 {
assert_eq!(result.unwrap(), None);
} else {
assert!(matches!(result, Err(Error::Status(actual)) if actual == status));
}
drop(failing);
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("recovered"))
.expect(1)
.mount(&server)
.await;
for _ in 0..2 {
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"recovered"
);
}
}
#[rstest]
#[tokio::test]
async fn failed_authentication_is_not_cached_and_does_not_read_secret() {
let server = MockServer::start().await;
let failing = Mock::given(path("/authn/acct/admin/authenticate"))
.respond_with(ResponseTemplate::new(401))
.expect(1)
.mount_as_scoped(&server)
.await;
let unused_secret = Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(0)
.mount_as_scoped(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager.async_read_secret("key").await,
Err(Error::AuthStatus(401))
));
drop(unused_secret);
drop(failing);
mount_auth(&server, 1).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(1)
.mount(&server)
.await;
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest]
#[tokio::test]
async fn trait_read_applies_cyberark_operation_timeout_to_authentication() {
let server = MockServer::start().await;
Mock::given(path("/authn/acct/admin/authenticate"))
.respond_with(
ResponseTemplate::new(200)
.set_body_string(TOKEN_JSON)
.set_delay(Duration::from_millis(50)),
)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
let context = SecretOperationContext::Cyberark(CyberarkOperationContext {
timeout: Some(Duration::from_millis(10)),
});
let result = BaseSecretManager::async_read_secret(&manager, "key", &context).await;
assert!(matches!(result, Err(Error::Http(error)) if error.is_timeout()));
}
#[rstest]
#[tokio::test]
async fn expired_tokens_and_secrets_are_fetched_again() {
let server = MockServer::start().await;
mount_auth(&server, 2).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_millis(1));
for _ in 0..2 {
assert!(manager.async_read_secret("key").await.unwrap().is_some());
tokio::time::sleep(Duration::from_millis(5)).await;
}
}
#[rstest]
#[case::plain("OPENAI_API_KEY")]
#[case::path("team/app/key")]
#[case::punctuation("a b+c.d-e_f~g")]
#[case::quote("needs \"quote\"")]
#[tokio::test]
async fn secret_names_use_python_quote_encoding(parity_fixture: ParityFixture, #[case] name: &str) {
let secret = parity_fixture
.secrets
.iter()
.find(|secret| secret.name == name)
.unwrap();
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(RawPath(secret.path.clone()))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(1)
.mount(&server)
.await;
assert_eq!(
manager(&server, Duration::from_secs(60))
.async_read_secret(name)
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest]
#[case::created(201)]
#[case::already_exists(409)]
#[case::unprocessable(422)]
#[case::server_error(500)]
#[tokio::test]
async fn writes_tolerate_policy_status_and_cache_value(#[case] policy_status: u16) {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/policies/acct/policy/root"))
.and(header("content-type", "application/x-yaml"))
.and(body_string("- !variable \"team/app\"\n"))
.respond_with(ResponseTemplate::new(policy_status))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/team%2Fapp"))
.and(body_string("v"))
.respond_with(ResponseTemplate::new(200))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
manager
.async_write_secret("team/app", &SecretValue::new("v"), None)
.await
.unwrap();
assert_eq!(
manager
.async_read_secret("team/app")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
}
#[rstest]
#[tokio::test]
async fn failed_value_write_is_not_cached() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/policies/acct/policy/root"))
.respond_with(ResponseTemplate::new(409))
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.and(body_string("v"))
.respond_with(ResponseTemplate::new(403))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("recovered"))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager
.async_write_secret("key", &SecretValue::new("v"), None)
.await,
Err(Error::Status(403))
));
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"recovered"
);
}
#[rstest]
#[case::parent("../etc")]
#[case::embedded_parent("team/../etc")]
#[case::control("key\n")]
#[tokio::test]
async fn unsafe_names_fail_before_http_calls(#[case] name: &str) {
let server = MockServer::start().await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager
.async_write_secret(name, &SecretValue::new("v"), None)
.await,
Err(Error::Operation(
litellm_secrets_types::Error::UnsafeSecretName
))
));
}
#[rstest]
#[tokio::test]
async fn delete_invalidates_cache_and_reports_not_supported() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("v"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
assert_eq!(
manager.async_delete_secret("key", Some(7)).await.unwrap(),
DeleteOutcome::NotSupported
);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
}
#[rstest]
fn new_validates_credentials_before_license_and_configuration() {
let empty: Arc<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
Arc::new(|_: &str| None);
assert!(matches!(
CyberArkSecretManager::new(empty, true),
Err(Error::MissingCredentials)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| (name == "CYBERARK_API_KEY").then(|| "k3y".into())),
false
),
Err(Error::EnterpriseRequired)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| (name == "CYBERARK_CLIENT_CERT").then(|| "cert".into())),
true
),
Err(Error::MissingCredentials)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_REFRESH_INTERVAL" => Some("abc".into()),
_ => None,
}),
true
),
Err(Error::RefreshInterval)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_API_BASE" => Some("not a url".into()),
_ => None,
}),
true
),
Err(Error::Endpoint)
));
}
#[rstest]
#[tokio::test]
async fn new_reads_environment_defaults_end_to_end() {
let server = MockServer::start().await;
Mock::given(path("/authn/default/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.mount(&server)
.await;
Mock::given(path("/secrets/default/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let endpoint = server.uri();
let manager = CyberArkSecretManager::new(
Arc::new(move |name: &str| match name {
"CYBERARK_API_BASE" => Some(endpoint.clone()),
"CYBERARK_API_KEY" => Some("k3y".into()),
_ => None,
}),
true,
)
.unwrap();
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest]
fn new_reports_missing_client_certificate_files() {
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()),
"CYBERARK_CLIENT_KEY" => Some("/missing/key".into()),
_ => None,
}),
true
),
Err(Error::ClientCertificate)
));
}
#[rstest]
#[tokio::test]
async fn trailing_slash_endpoint_preserves_base_path() {
let server = MockServer::start().await;
Mock::given(path("/prefix/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/prefix/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let endpoint = format!("{}/prefix/", server.uri()).parse().unwrap();
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
endpoint,
"acct".into(),
"admin".into(),
SecretValue::new("k3y"),
Some(Duration::from_secs(60)),
);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest]
fn parity_fixture_matches_authentication_contract(parity_fixture: ParityFixture) {
assert_eq!(parity_fixture.endpoint, "http://conjur.test:8080");
assert_eq!(parity_fixture.account, "acct");
assert_eq!(parity_fixture.username, "admin");
assert_eq!(parity_fixture.api_key, "k3y");
assert_eq!(
parity_fixture.authenticate_path,
"/authn/acct/admin/authenticate"
);
assert_eq!(parity_fixture.token_json, TOKEN_JSON);
assert_eq!(
parity_fixture.authorization_header,
format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON))
);
assert_eq!(parity_fixture.policy_path, "/policies/acct/policy/root");
assert_eq!(parity_fixture.secrets.len(), 4);
assert_eq!(
parity_fixture.secrets[1].policy_body,
"- !variable \"team/app/key\"\n"
);
}
#[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,452 @@
use super::*;
#[rstest]
#[tokio::test]
async fn successful_reads_cache_auth_secret_and_redact_values() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
let token = STANDARD.encode(TOKEN_JSON);
Mock::given(path("/secrets/acct/variable/OPENAI_API_KEY"))
.and(header("authorization", format!("Token token=\"{token}\"")))
.respond_with(ResponseTemplate::new(200).set_body_string("sk-live"))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
for _ in 0..2 {
let value = manager
.async_read_secret("OPENAI_API_KEY")
.await
.unwrap()
.unwrap();
assert_eq!(value.expose(), "sk-live");
assert!(!format!("{value:?}").contains("sk-live"));
}
}
#[rstest]
#[tokio::test]
async fn concurrent_reads_share_authentication_and_secret_requests() {
let server = MockServer::start().await;
Mock::given(path("/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(
ResponseTemplate::new(200)
.set_body_string(TOKEN_JSON)
.set_delay(Duration::from_millis(20)),
)
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.and(header(
"authorization",
format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON)),
))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
let (first, second) = tokio::join!(
manager.async_read_secret("key"),
manager.async_read_secret("key")
);
assert_eq!(first.unwrap().unwrap().expose(), "value");
assert_eq!(second.unwrap().unwrap().expose(), "value");
}
#[rstest]
#[case::host("host/team/app", "/authn/acct/host%2Fteam%2Fapp/authenticate")]
#[case::user("alice@devops", "/authn/acct/alice%40devops/authenticate")]
#[tokio::test]
async fn authentication_encodes_login(#[case] username: &str, #[case] expected_path: &str) {
let server = MockServer::start().await;
Mock::given(RawPath(expected_path.to_owned()))
.and(method("POST"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
server.uri().parse().unwrap(),
"acct".into(),
username.into(),
SecretValue::new("k3y"),
None,
);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest]
#[tokio::test]
async fn rejected_cached_token_is_reauthenticated_once() {
let server = MockServer::start().await;
mount_auth(&server, 2).await;
Mock::given(path("/secrets/acct/variable/first"))
.respond_with(ResponseTemplate::new(200).set_body_string("first-value"))
.expect(1)
.mount(&server)
.await;
let attempts = Arc::new(AtomicUsize::new(0));
let attempts_for_response = Arc::clone(&attempts);
Mock::given(path("/secrets/acct/variable/second"))
.respond_with(move |_: &Request| {
if attempts_for_response.fetch_add(1, Ordering::SeqCst) == 0 {
ResponseTemplate::new(401)
} else {
ResponseTemplate::new(200).set_body_string("second-value")
}
})
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(600));
assert_eq!(
manager
.async_read_secret("first")
.await
.unwrap()
.unwrap()
.expose(),
"first-value"
);
assert_eq!(
manager
.async_read_secret("second")
.await
.unwrap()
.unwrap()
.expose(),
"second-value"
);
assert_eq!(attempts.load(Ordering::SeqCst), 2);
}
#[rstest]
#[tokio::test]
async fn failed_authentication_is_not_cached_and_does_not_read_secret() {
let server = MockServer::start().await;
let failing = Mock::given(path("/authn/acct/admin/authenticate"))
.respond_with(ResponseTemplate::new(401))
.expect(1)
.mount_as_scoped(&server)
.await;
let unused_secret = Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(0)
.mount_as_scoped(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager.async_read_secret("key").await,
Err(Error::AuthStatus(401))
));
drop(unused_secret);
drop(failing);
mount_auth(&server, 1).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(1)
.mount(&server)
.await;
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest]
#[tokio::test]
async fn trait_read_applies_cyberark_operation_timeout_to_authentication() {
let server = MockServer::start().await;
Mock::given(path("/authn/acct/admin/authenticate"))
.respond_with(
ResponseTemplate::new(200)
.set_body_string(TOKEN_JSON)
.set_delay(Duration::from_millis(50)),
)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
let context = CyberarkOperationContext {
timeout: Some(Duration::from_millis(10)),
};
let result = BaseSecretManager::async_read_secret(&manager, "key", &context).await;
match result {
Err(Error::Timeout) => {}
Err(Error::Http(error)) => assert!(error.is_timeout()),
other => panic!("expected timeout, got {other:?}"),
}
}
#[rstest]
fn new_validates_credentials_before_license_and_configuration() {
let empty: Arc<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
Arc::new(|_: &str| None);
assert!(matches!(
CyberArkSecretManager::new(empty, true),
Err(Error::MissingCredentials)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| (name == "CYBERARK_API_KEY").then(|| "k3y".into())),
false
),
Err(Error::EnterpriseRequired)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| (name == "CYBERARK_CLIENT_CERT").then(|| "cert".into())),
true
),
Err(Error::MissingCredentials)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_REFRESH_INTERVAL" => Some("abc".into()),
_ => None,
}),
true
),
Err(Error::RefreshInterval)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_API_BASE" => Some("not a url".into()),
_ => None,
}),
true
),
Err(Error::Endpoint)
));
}
#[rstest]
fn certificate_only_credentials_are_validated_as_a_client_identity() {
let result = CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()),
"CYBERARK_CLIENT_KEY" => Some("/missing/key".into()),
_ => None,
}),
true,
);
assert!(matches!(result, Err(Error::ClientCertificate)));
}
#[rstest]
#[case::certificate_only("")]
#[case::certificate_and_api_key("k3y")]
#[tokio::test]
async fn configured_client_identity_preserves_auth_request_and_read_result(
client_identity_directory: tempfile::TempDir,
#[case] api_key: &'static str,
) {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/authn/default/admin/authenticate"))
.and(body_string(api_key))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.expect(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/secrets/default/variable/key"))
.and(header(
"authorization",
format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON)),
))
.respond_with(ResponseTemplate::new(200).set_body_string(" value\n"))
.expect(1)
.mount(&server)
.await;
let endpoint = server.uri();
let certificate = client_identity_directory.path().join("client.crt");
let key = client_identity_directory.path().join("client.key");
let manager = CyberArkSecretManager::new(
Arc::new(move |name: &str| match name {
"CYBERARK_API_BASE" => Some(endpoint.clone()),
"CYBERARK_API_KEY" => Some(api_key.into()),
"CYBERARK_CLIENT_CERT" => Some(certificate.to_str().unwrap().into()),
"CYBERARK_CLIENT_KEY" => Some(key.to_str().unwrap().into()),
_ => None,
}),
true,
)
.unwrap();
assert!(server.received_requests().await.unwrap().is_empty());
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
" value\n"
);
}
#[rstest]
#[case::certificate_only("", "client.crt")]
#[case::key_only("", "client.key")]
#[case::certificate_with_api_key("k3y", "client.crt")]
#[case::key_with_api_key("k3y", "client.key")]
fn invalid_client_identity_is_not_ignored(
client_identity_directory: tempfile::TempDir,
#[case] api_key: &'static str,
#[case] invalid_file: &str,
) {
std::fs::write(
client_identity_directory.path().join(invalid_file),
"not PEM",
)
.unwrap();
let certificate = client_identity_directory.path().join("client.crt");
let key = client_identity_directory.path().join("client.key");
let result = CyberArkSecretManager::new(
Arc::new(move |name: &str| match name {
"CYBERARK_API_KEY" => Some(api_key.into()),
"CYBERARK_CLIENT_CERT" => Some(certificate.to_str().unwrap().into()),
"CYBERARK_CLIENT_KEY" => Some(key.to_str().unwrap().into()),
_ => None,
}),
true,
);
assert!(matches!(result, Err(Error::ClientCertificate)));
}
#[rstest]
#[case::certificate_only("")]
#[case::certificate_and_api_key("k3y")]
fn client_identity_does_not_bypass_the_enterprise_requirement(#[case] api_key: &'static str) {
let result = CyberArkSecretManager::new(
Arc::new(move |name: &str| match name {
"CYBERARK_API_KEY" => Some(api_key.into()),
"CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()),
"CYBERARK_CLIENT_KEY" => Some("/missing/key".into()),
_ => None,
}),
false,
);
assert!(matches!(result, Err(Error::EnterpriseRequired)));
}
#[rstest]
#[tokio::test]
async fn new_reads_environment_defaults_end_to_end() {
let server = MockServer::start().await;
Mock::given(path("/authn/default/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.mount(&server)
.await;
Mock::given(path("/secrets/default/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let endpoint = server.uri();
let manager = CyberArkSecretManager::new(
Arc::new(move |name: &str| match name {
"CYBERARK_API_BASE" => Some(endpoint.clone()),
"CYBERARK_API_KEY" => Some("k3y".into()),
_ => None,
}),
true,
)
.unwrap();
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest]
fn new_reports_missing_client_certificate_files() {
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()),
"CYBERARK_CLIENT_KEY" => Some("/missing/key".into()),
_ => None,
}),
true
),
Err(Error::ClientCertificate)
));
}
#[rstest]
#[tokio::test]
async fn trailing_slash_endpoint_preserves_base_path() {
let server = MockServer::start().await;
Mock::given(path("/prefix/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/prefix/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let endpoint = format!("{}/prefix/", server.uri()).parse().unwrap();
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
endpoint,
"acct".into(),
"admin".into(),
SecretValue::new("k3y"),
Some(Duration::from_secs(60)),
);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}

View file

@ -0,0 +1,145 @@
use super::*;
#[rstest]
#[tokio::test]
async fn a_rejected_refreshed_token_surfaces_the_error_without_another_retry() {
let server = MockServer::start().await;
mount_auth(&server, 2).await;
Mock::given(path("/secrets/acct/variable/first"))
.respond_with(ResponseTemplate::new(200).set_body_string("first-value"))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/second"))
.respond_with(ResponseTemplate::new(401))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert_eq!(
manager
.async_read_secret("first")
.await
.unwrap()
.unwrap()
.expose(),
"first-value"
);
let result = tokio::time::timeout(Duration::from_secs(5), manager.async_read_secret("second"))
.await
.expect("authentication retries must terminate");
assert!(matches!(result, Err(Error::Status(401))));
}
#[rstest]
#[case::not_found(404)]
#[case::unauthorized(401)]
#[case::forbidden(403)]
#[case::server_error(500)]
#[tokio::test]
async fn failed_reads_are_not_cached(#[case] status: u16) {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
let failing = Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(status))
.expect(1)
.mount_as_scoped(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
let result = manager.async_read_secret("key").await;
if status == 404 {
assert_eq!(result.unwrap(), None);
} else {
assert!(matches!(result, Err(Error::Status(actual)) if actual == status));
}
drop(failing);
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("recovered"))
.expect(1)
.mount(&server)
.await;
for _ in 0..2 {
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"recovered"
);
}
}
#[rstest]
#[tokio::test]
async fn expired_tokens_and_secrets_are_fetched_again() {
let server = MockServer::start().await;
mount_auth(&server, 2).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_millis(1));
for _ in 0..2 {
assert!(manager.async_read_secret("key").await.unwrap().is_some());
tokio::time::sleep(Duration::from_millis(5)).await;
}
}
#[rstest]
#[case::plain("OPENAI_API_KEY")]
#[case::path("team/app/key")]
#[case::punctuation("a b+c.d-e_f~g")]
#[case::quote("needs \"quote\"")]
#[tokio::test]
async fn secret_names_use_python_quote_encoding(parity_fixture: ParityFixture, #[case] name: &str) {
let secret = parity_fixture
.secrets
.iter()
.find(|secret| secret.name == name)
.unwrap();
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(RawPath(secret.path.clone()))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(1)
.mount(&server)
.await;
assert_eq!(
manager(&server, Duration::from_secs(60))
.async_read_secret(name)
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest]
#[case::parent("../etc")]
#[case::embedded_parent("team/../etc")]
#[case::control("key\n")]
#[tokio::test]
async fn unsafe_names_fail_before_http_calls(#[case] name: &str) {
let server = MockServer::start().await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager.async_read_secret(name).await,
Err(Error::Operation(
litellm_secrets_types::Error::UnsafeSecretName
))
));
assert!(matches!(
manager
.async_write_secret(name, &SecretValue::new("v"), None)
.await,
Err(Error::Operation(
litellm_secrets_types::Error::UnsafeSecretName
))
));
}

View file

@ -0,0 +1,70 @@
use super::*;
pub(super) const TOKEN_JSON: &str = r#"{"protected":"p","payload":"q","signature":"s"}"#;
#[derive(Deserialize)]
pub(super) struct ParityFixture {
pub(super) account: String,
pub(super) username: String,
pub(super) api_key: String,
pub(super) authenticate_path: String,
pub(super) token_json: String,
pub(super) authorization_header: String,
pub(super) policy_path: String,
pub(super) secrets: Vec<ParitySecret>,
}
#[derive(Deserialize)]
pub(super) struct ParitySecret {
pub(super) name: String,
pub(super) path: String,
pub(super) policy_body: String,
}
#[derive(Debug)]
pub(super) struct RawPath(pub(super) String);
impl Match for RawPath {
fn matches(&self, request: &Request) -> bool {
request.url.path() == self.0
}
}
#[fixture]
pub(super) fn parity_fixture() -> ParityFixture {
serde_json::from_str(include_str!("../fixtures/parity.json")).unwrap()
}
#[fixture]
pub(super) fn client_identity_directory() -> tempfile::TempDir {
let identity = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap();
let directory = tempfile::tempdir().unwrap();
std::fs::write(directory.path().join("client.crt"), identity.cert.pem()).unwrap();
std::fs::write(
directory.path().join("client.key"),
identity.signing_key.serialize_pem(),
)
.unwrap();
directory
}
pub(super) fn manager(server: &MockServer, ttl: Duration) -> CyberArkSecretManager {
CyberArkSecretManager::with_client(
reqwest::Client::new(),
server.uri().parse().unwrap(),
"acct".into(),
"admin".into(),
SecretValue::new("k3y"),
Some(ttl),
)
}
pub(super) async fn mount_auth(server: &MockServer, expected: u64) {
Mock::given(method("POST"))
.and(path("/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.expect(expected)
.mount(server)
.await;
}

View file

@ -0,0 +1,450 @@
use super::*;
#[rstest]
#[tokio::test]
async fn rejected_write_token_is_reauthenticated_once() {
let server = MockServer::start().await;
mount_auth(&server, 2).await;
Mock::given(path("/policies/acct/policy/root"))
.respond_with(ResponseTemplate::new(409))
.expect(1)
.mount(&server)
.await;
let attempts = Arc::new(AtomicUsize::new(0));
let attempts_for_response = Arc::clone(&attempts);
Mock::given(method("POST"))
.and(path("/secrets/acct/variable/key"))
.and(body_string("value"))
.respond_with(move |_: &Request| {
if attempts_for_response.fetch_add(1, Ordering::SeqCst) == 0 {
ResponseTemplate::new(401)
} else {
ResponseTemplate::new(200)
}
})
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(600));
manager
.async_write_secret("key", &SecretValue::new("value"), None)
.await
.unwrap();
assert_eq!(attempts.load(Ordering::SeqCst), 2);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest]
#[case::created(201)]
#[case::already_exists(409)]
#[case::unprocessable(422)]
#[case::server_error(500)]
#[tokio::test]
async fn writes_tolerate_policy_status_and_cache_value(#[case] policy_status: u16) {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/policies/acct/policy/root"))
.and(header("content-type", "application/x-yaml"))
.and(body_string("- !variable \"team/app\"\n"))
.respond_with(ResponseTemplate::new(policy_status))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/team%2Fapp"))
.and(body_string("v"))
.respond_with(ResponseTemplate::new(200))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
manager
.async_write_secret("team/app", &SecretValue::new("v"), None)
.await
.unwrap();
assert_eq!(
manager
.async_read_secret("team/app")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
}
#[rstest]
#[tokio::test]
async fn failed_value_write_is_not_cached() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/policies/acct/policy/root"))
.respond_with(ResponseTemplate::new(409))
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.and(body_string("v"))
.respond_with(ResponseTemplate::new(403))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("recovered"))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager
.async_write_secret("key", &SecretValue::new("v"), None)
.await,
Err(Error::Status(403))
));
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"recovered"
);
}
#[rstest]
#[tokio::test]
async fn delete_invalidates_cache_and_reports_not_supported() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("v"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
assert_eq!(
manager.async_delete_secret("key", Some(7)).await.unwrap(),
DeleteOutcome::NotSupported
);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
}
#[rstest]
#[tokio::test]
async fn writes_match_python_parity_fixture(parity_fixture: ParityFixture) {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path(&parity_fixture.authenticate_path))
.and(body_string(&parity_fixture.api_key))
.respond_with(ResponseTemplate::new(200).set_body_string(&parity_fixture.token_json))
.expect(1)
.mount(&server)
.await;
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
server.uri().parse().unwrap(),
parity_fixture.account,
parity_fixture.username,
SecretValue::new(parity_fixture.api_key),
Some(Duration::from_secs(60)),
);
for secret in parity_fixture.secrets {
Mock::given(method("POST"))
.and(path(&parity_fixture.policy_path))
.and(header(
"authorization",
&parity_fixture.authorization_header,
))
.and(header("content-type", "application/x-yaml"))
.and(body_string(&secret.policy_body))
.respond_with(ResponseTemplate::new(201))
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(RawPath(secret.path))
.and(header(
"authorization",
&parity_fixture.authorization_header,
))
.and(body_string("value"))
.respond_with(ResponseTemplate::new(201))
.expect(1)
.mount(&server)
.await;
manager
.async_write_secret(&secret.name, &SecretValue::new("value"), None)
.await
.unwrap();
assert_eq!(
manager
.async_read_secret(&secret.name)
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
}
#[rstest]
#[tokio::test]
#[ignore]
async fn live_conjur_round_trip() {
let endpoint: reqwest::Url = std::env::var("CYBERARK_API_BASE").unwrap().parse().unwrap();
let account = std::env::var("CYBERARK_ACCOUNT").unwrap();
let username = std::env::var("CYBERARK_USERNAME").unwrap();
let api_key = SecretValue::new(std::env::var("CYBERARK_API_KEY").unwrap());
let name = format!(
"{}-{}",
std::env::var("LITELLM_CONJUR_LIVE_SECRET_NAME").unwrap(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
);
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
endpoint.clone(),
account.clone(),
username.clone(),
api_key.clone(),
Some(Duration::from_secs(60)),
);
assert!(manager.async_read_secret(&name).await.unwrap().is_none());
for expected in ["first-π\n", " second-π "] {
manager
.async_write_secret(&name, &SecretValue::new(expected), None)
.await
.unwrap();
let verifier = CyberArkSecretManager::with_client(
reqwest::Client::new(),
endpoint.clone(),
account.clone(),
username.clone(),
api_key.clone(),
Some(Duration::from_secs(60)),
);
assert_eq!(
verifier
.async_read_secret(&name)
.await
.unwrap()
.unwrap()
.expose(),
expected
);
}
manager
.async_rotate_secret(&name, &name, &SecretValue::new("rotated-value"))
.await
.unwrap();
let alias = format!("{name}-rotated");
manager
.async_rotate_secret(&name, &alias, &SecretValue::new("new-alias-value"))
.await
.unwrap();
assert_eq!(
manager
.async_read_secret(&name)
.await
.unwrap()
.unwrap()
.expose(),
"rotated-value"
);
assert_eq!(
manager
.async_read_secret(&alias)
.await
.unwrap()
.unwrap()
.expose(),
"new-alias-value"
);
}
#[rstest]
#[case::same_alias("old")]
#[case::new_alias("new")]
#[tokio::test]
async fn rotation_stores_the_replacement_and_retains_other_aliases(#[case] new_name: &'static str) {
use std::sync::atomic::{AtomicBool, Ordering};
let server = MockServer::start().await;
mount_auth(&server, 1).await;
let written = Arc::new(AtomicBool::new(false));
let read_state = written.clone();
Mock::given(method("GET"))
.respond_with(move |request: &wiremock::Request| {
let name = request.url.path().rsplit('/').next().unwrap();
let value = if name == new_name && read_state.load(Ordering::SeqCst) {
"new-value"
} else {
"old-value"
};
ResponseTemplate::new(200).set_body_string(value)
})
.expect(if new_name == "old" { 2 } else { 3 })
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/policies/acct/policy/root"))
.and(body_string(format!("- !variable \"{new_name}\"\n")))
.respond_with(ResponseTemplate::new(201))
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path(format!("/secrets/acct/variable/{new_name}")))
.and(body_string("new-value"))
.respond_with(move |_: &wiremock::Request| {
written.store(true, Ordering::SeqCst);
ResponseTemplate::new(201)
})
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
manager
.async_rotate_secret("old", new_name, &SecretValue::new("new-value"))
.await
.unwrap();
assert_eq!(
manager
.async_read_secret(new_name)
.await
.unwrap()
.unwrap()
.expose(),
"new-value"
);
assert_eq!(
manager
.async_read_secret("old")
.await
.unwrap()
.unwrap()
.expose(),
if new_name == "old" {
"new-value"
} else {
"old-value"
}
);
assert!(
server
.received_requests()
.await
.unwrap()
.iter()
.all(|request| request.method != "DELETE")
);
}
#[tokio::test]
async fn rotation_verifies_the_remote_value_instead_of_the_write_cache() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(method("GET"))
.respond_with(ResponseTemplate::new(200).set_body_string("unchanged"))
.expect(2)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/policies/acct/policy/root"))
.respond_with(ResponseTemplate::new(201))
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/secrets/acct/variable/new"))
.respond_with(ResponseTemplate::new(201))
.expect(1)
.mount(&server)
.await;
let result = manager(&server, Duration::from_secs(60))
.async_rotate_secret("old", "new", &SecretValue::new("replacement"))
.await;
assert!(matches!(
result,
Err(litellm_secrets_types::RotationError::Verification {
source: Error::Operation(litellm_secrets_types::Error::NewSecretMismatch),
..
})
));
}
#[rstest]
#[case::colon("foo: bar")]
#[case::comment("foo # bar")]
#[case::plain("plain-alias")]
#[case::email("team/user@example.com")]
#[case::quote("needs \"quote\"")]
#[case::backslash("a\\b")]
#[tokio::test]
async fn policy_writes_preserve_yaml_metacharacters_as_one_variable(#[case] name: &'static str) {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(method("POST"))
.and(path("/policies/acct/policy/root"))
.respond_with(move |request: &wiremock::Request| {
let body = std::str::from_utf8(&request.body).unwrap();
let scalar = body
.strip_prefix("- !variable ")
.unwrap()
.strip_suffix('\n')
.unwrap();
assert_eq!(serde_json::from_str::<String>(scalar).unwrap(), name);
ResponseTemplate::new(201)
})
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(body_string("value"))
.respond_with(ResponseTemplate::new(201))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
manager
.async_write_secret(name, &SecretValue::new("value"), None)
.await
.unwrap();
assert_eq!(
manager
.async_read_secret(name)
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}

View file

@ -0,0 +1 @@
- https://docs.cloud.google.com/secret-manager/docs/reference/rest/v1/projects.secrets.versions/access

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