mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Merge remote-tracking branch 'origin/main' into litellm_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
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:
commit
09cc332c03
346 changed files with 32327 additions and 7027 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
||||
|
|
|
|||
|
|
@ -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
15
.github/merge-smoke-tests.json
vendored
Normal 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"
|
||||
}
|
||||
}
|
||||
58
.github/scripts/assert_ci_coverage.py
vendored
58
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -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
493
.github/scripts/run_merge_smoke.py
vendored
Normal 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())
|
||||
2
.github/scripts/verify_linux_native_wheel.py
vendored
2
.github/scripts/verify_linux_native_wheel.py
vendored
|
|
@ -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),
|
||||
|
|
|
|||
2
.github/workflows/check-ui-api-types.yml
vendored
2
.github/workflows/check-ui-api-types.yml
vendored
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
|
|
|
|||
3
.github/workflows/ci-coverage.yml
vendored
3
.github/workflows/ci-coverage.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
8
.github/workflows/codeql.yml
vendored
8
.github/workflows/codeql.yml
vendored
|
|
@ -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
|
||||
|
||||
|
|
|
|||
2
.github/workflows/codspeed.yml
vendored
2
.github/workflows/codspeed.yml
vendored
|
|
@ -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/**"
|
||||
|
|
|
|||
2
.github/workflows/cost-map-guard.yml
vendored
2
.github/workflows/cost-map-guard.yml
vendored
|
|
@ -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:
|
||||
|
|
|
|||
186
.github/workflows/create-release.yml
vendored
186
.github/workflows/create-release.yml
vendored
|
|
@ -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 }}
|
||||
|
|
@ -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"
|
||||
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "uv.lock"
|
||||
|
|
|
|||
1
.github/workflows/image-scan.yml
vendored
1
.github/workflows/image-scan.yml
vendored
|
|
@ -4,7 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_branch
|
||||
- "litellm_**"
|
||||
paths:
|
||||
|
|
|
|||
2
.github/workflows/osv-scan.yml
vendored
2
.github/workflows/osv-scan.yml
vendored
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
schedule:
|
||||
- cron: "23 6 * * *"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
6
.github/workflows/test-code-quality.yml
vendored
6
.github/workflows/test-code-quality.yml
vendored
|
|
@ -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
|
||||
|
||||
|
|
|
|||
2
.github/workflows/test-linting.yml
vendored
2
.github/workflows/test-linting.yml
vendored
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
|
|
|
|||
2
.github/workflows/test-litellm-ui-build.yml
vendored
2
.github/workflows/test-litellm-ui-build.yml
vendored
|
|
@ -7,8 +7,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
concurrency:
|
||||
|
|
|
|||
2
.github/workflows/test-litellm-ui-lint.yml
vendored
2
.github/workflows/test-litellm-ui-lint.yml
vendored
|
|
@ -6,8 +6,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
concurrency:
|
||||
|
|
|
|||
3
.github/workflows/test-litellm-ui-unit.yml
vendored
3
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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
95
.github/workflows/test-merge-smoke.yml
vendored
Normal 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
|
||||
3
.github/workflows/test-postgres.yml
vendored
3
.github/workflows/test-postgres.yml
vendored
|
|
@ -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:
|
||||
|
|
|
|||
2
.github/workflows/test-redis-compat.yml
vendored
2
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "litellm/_redis.py"
|
||||
|
|
|
|||
2
.github/workflows/test-rust.yml
vendored
2
.github/workflows/test-rust.yml
vendored
|
|
@ -29,8 +29,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "litellm-rust/**"
|
||||
|
|
|
|||
2
.github/workflows/test-semgrep.yml
vendored
2
.github/workflows/test-semgrep.yml
vendored
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
|
|
|
|||
2
.github/workflows/test-terraform-modules.yml
vendored
2
.github/workflows/test-terraform-modules.yml
vendored
|
|
@ -9,8 +9,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "terraform/litellm/aws/**"
|
||||
|
|
|
|||
|
|
@ -8,8 +8,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "terraform/provider/**"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
3
.github/workflows/test-unit-proxy-db.yml
vendored
3
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
3
.github/workflows/test-unit.yml
vendored
3
.github/workflows/test-unit.yml
vendored
|
|
@ -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:
|
||||
|
|
|
|||
2
.github/workflows/test-vscode-extension.yml
vendored
2
.github/workflows/test-vscode-extension.yml
vendored
|
|
@ -6,8 +6,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "vscode-extension/**"
|
||||
|
|
|
|||
4
.github/workflows/zizmor.yml
vendored
4
.github/workflows/zizmor.yml
vendored
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
193
litellm-rust/Cargo.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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!({})))
|
||||
|
|
|
|||
|
|
@ -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`.
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
pub mod secrets;
|
||||
|
|
@ -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) })
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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()),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>(),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
);
|
||||
});
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
113
litellm-rust/crates/python-bridge/src/secrets/mutation.rs
Normal file
113
litellm-rust/crates/python-bridge/src/secrets/mutation.rs
Normal 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}")),
|
||||
}
|
||||
}
|
||||
247
litellm-rust/crates/python-bridge/src/secrets/operations.rs
Normal file
247
litellm-rust/crates/python-bridge/src/secrets/operations.rs
Normal 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)
|
||||
}
|
||||
150
litellm-rust/crates/python-bridge/src/secrets/provider.rs
Normal file
150
litellm-rust/crates/python-bridge/src/secrets/provider.rs
Normal 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)
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
463
litellm-rust/crates/python-bridge/src/secrets/runtime.rs
Normal file
463
litellm-rust/crates/python-bridge/src/secrets/runtime.rs
Normal 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(¤t, "__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,
|
||||
¤t_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)
|
||||
}
|
||||
182
litellm-rust/crates/python-bridge/src/secrets/vault.rs
Normal file
182
litellm-rust/crates/python-bridge/src/secrets/vault.rs
Normal 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()?,
|
||||
})
|
||||
}
|
||||
}
|
||||
168
litellm-rust/crates/python-bridge/src/secrets/vault/operation.rs
Normal file
168
litellm-rust/crates/python-bridge/src/secrets/vault/operation.rs
Normal 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())
|
||||
}
|
||||
1
litellm-rust/crates/secrets-aws/AGENTS.md
Normal file
1
litellm-rust/crates/secrets-aws/AGENTS.md
Normal file
|
|
@ -0,0 +1 @@
|
|||
- https://docs.aws.amazon.com/secretsmanager/latest/apireference/Welcome.html
|
||||
|
|
@ -22,3 +22,4 @@ base64.workspace = true
|
|||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
tempfile = "3"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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())?))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
}
|
||||
|
|
|
|||
116
litellm-rust/crates/secrets-aws/src/secret_manager/client.rs
Normal file
116
litellm-rust/crates/secrets-aws/src/secret_manager/client.rs
Normal 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))
|
||||
}
|
||||
}
|
||||
225
litellm-rust/crates/secrets-aws/src/secret_manager/read.rs
Normal file
225
litellm-rust/crates/secrets-aws/src/secret_manager/read.rs
Normal 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))
|
||||
}
|
||||
}
|
||||
}
|
||||
329
litellm-rust/crates/secrets-aws/src/secret_manager/write.rs
Normal file
329
litellm-rust/crates/secrets-aws/src/secret_manager/write.rs
Normal 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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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()
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
);
|
||||
}
|
||||
198
litellm-rust/crates/secrets-aws/tests/secret_manager/reads.rs
Normal file
198
litellm-rust/crates/secrets-aws/tests/secret_manager/reads.rs
Normal 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(_))
|
||||
));
|
||||
}
|
||||
|
|
@ -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())
|
||||
}
|
||||
615
litellm-rust/crates/secrets-aws/tests/secret_manager/writes.rs
Normal file
615
litellm-rust/crates/secrets-aws/tests/secret_manager/writes.rs
Normal 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(_)))
|
||||
}
|
||||
}
|
||||
}
|
||||
1
litellm-rust/crates/secrets-azure/AGENTS.md
Normal file
1
litellm-rust/crates/secrets-azure/AGENTS.md
Normal file
|
|
@ -0,0 +1 @@
|
|||
- https://learn.microsoft.com/en-us/rest/api/keyvault/secrets/get-secret/get-secret
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
));
|
||||
}
|
||||
|
|
|
|||
1
litellm-rust/crates/secrets-cyberark/AGENTS.md
Normal file
1
litellm-rust/crates/secrets-cyberark/AGENTS.md
Normal file
|
|
@ -0,0 +1 @@
|
|||
- https://docs.cyberark.com/conjur-open-source/latest/en/content/developer/conjur_api_retrieve_secret.htm
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
118
litellm-rust/crates/secrets-cyberark/src/secret_manager/read.rs
Normal file
118
litellm-rust/crates/secrets-cyberark/src/secret_manager/read.rs
Normal 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
|
||||
}
|
||||
}
|
||||
262
litellm-rust/crates/secrets-cyberark/src/secret_manager/write.rs
Normal file
262
litellm-rust/crates/secrets-cyberark/src/secret_manager/write.rs
Normal 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
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
);
|
||||
}
|
||||
|
|
@ -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
|
||||
))
|
||||
));
|
||||
}
|
||||
|
|
@ -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;
|
||||
}
|
||||
|
|
@ -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"
|
||||
);
|
||||
}
|
||||
1
litellm-rust/crates/secrets-google/AGENTS.md
Normal file
1
litellm-rust/crates/secrets-google/AGENTS.md
Normal 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
Loading…
Add table
Reference in a new issue