mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
Merge remote-tracking branch 'origin/main' into litellm_native_redis_semantic_cache
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
615a226a30
37 changed files with 3025 additions and 32 deletions
2
.github/workflows/test-rust.yml
vendored
2
.github/workflows/test-rust.yml
vendored
|
|
@ -130,7 +130,7 @@ jobs:
|
|||
- name: Test secret manager feature combinations
|
||||
run: |
|
||||
cargo test -p litellm-auth-gcp --locked --no-default-features
|
||||
for features in '' aws google azure cyberark aws,google aws,azure google,azure aws,google,azure aws,google,cyberark aws,google,azure,cyberark; do
|
||||
for features in '' aws google hashicorp azure cyberark aws,google aws,azure google,azure aws,google,azure aws,google,cyberark aws,google,azure,cyberark aws,google,hashicorp,azure,cyberark; do
|
||||
cargo test -p litellm-secrets --locked --no-default-features --features "$features"
|
||||
done
|
||||
|
||||
|
|
|
|||
124
litellm-rust/Cargo.lock
generated
124
litellm-rust/Cargo.lock
generated
|
|
@ -2611,6 +2611,21 @@ dependencies = [
|
|||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-cache-gcs"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"futures-util",
|
||||
"litellm-auth-gcp",
|
||||
"litellm-auth-types",
|
||||
"litellm-cache",
|
||||
"percent-encoding",
|
||||
"reqwest 0.12.28",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-cache-memory"
|
||||
version = "0.1.0"
|
||||
|
|
@ -2830,6 +2845,7 @@ dependencies = [
|
|||
"litellm-cache",
|
||||
"litellm-cache-azure-blob",
|
||||
"litellm-cache-disk",
|
||||
"litellm-cache-gcs",
|
||||
"litellm-cache-memory",
|
||||
"litellm-cache-redis",
|
||||
"litellm-cache-redis-semantic",
|
||||
|
|
@ -2866,6 +2882,7 @@ dependencies = [
|
|||
"litellm-secrets-azure",
|
||||
"litellm-secrets-cyberark",
|
||||
"litellm-secrets-google",
|
||||
"litellm-secrets-hashicorp",
|
||||
"litellm-secrets-types",
|
||||
"moka",
|
||||
"reqwest 0.12.28",
|
||||
|
|
@ -2963,6 +2980,26 @@ dependencies = [
|
|||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-secrets-hashicorp"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"litellm-core-utils",
|
||||
"litellm-secrets-types",
|
||||
"moka",
|
||||
"rstest",
|
||||
"rustify",
|
||||
"rustify_derive",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tempfile",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"vaultrs",
|
||||
"veil",
|
||||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-secrets-types"
|
||||
version = "0.1.0"
|
||||
|
|
@ -4201,6 +4238,40 @@ dependencies = [
|
|||
"semver",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustify"
|
||||
version = "0.7.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "4800ce4c1cc2fec12c559dae2ddbf0e17fcee7569b796e6d75898efef443368b"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"async-trait",
|
||||
"bytes",
|
||||
"http 1.4.2",
|
||||
"reqwest 0.13.5",
|
||||
"rustify_derive",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"serde_urlencoded",
|
||||
"thiserror 1.0.69",
|
||||
"tracing",
|
||||
"url",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustify_derive"
|
||||
version = "0.5.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "78ea7fda74240f7410d0198b603a8a2f662acc7d76b6667a49f9b162cd8d9b4f"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"regex",
|
||||
"serde_urlencoded",
|
||||
"syn 1.0.109",
|
||||
"synstructure 0.12.6",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustix"
|
||||
version = "1.1.5"
|
||||
|
|
@ -4753,6 +4824,17 @@ version = "2.6.1"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292"
|
||||
|
||||
[[package]]
|
||||
name = "syn"
|
||||
version = "1.0.109"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "72b64191b275b66ffe2469e8af2c1cfe3bafa67b529ead792a6d0160888b4237"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"unicode-ident",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "syn"
|
||||
version = "2.0.119"
|
||||
|
|
@ -4784,6 +4866,18 @@ dependencies = [
|
|||
"futures-core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "synstructure"
|
||||
version = "0.12.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f36bdaa60a83aca3921b5259d5400cbf5e90fc51931376a9bd4a0eb79aa7210f"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 1.0.109",
|
||||
"unicode-xid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "synstructure"
|
||||
version = "0.13.2"
|
||||
|
|
@ -5198,6 +5292,7 @@ version = "0.1.44"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100"
|
||||
dependencies = [
|
||||
"log",
|
||||
"pin-project-lite",
|
||||
"tracing-attributes",
|
||||
"tracing-core",
|
||||
|
|
@ -5382,6 +5477,12 @@ version = "1.13.3"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c6f5d3c3b1bf09027a88a6bc961fc00497d651009560b5463668dc81b0fa87a8"
|
||||
|
||||
[[package]]
|
||||
name = "unicode-xid"
|
||||
version = "0.2.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853"
|
||||
|
||||
[[package]]
|
||||
name = "unicode_categories"
|
||||
version = "0.1.1"
|
||||
|
|
@ -5441,6 +5542,25 @@ version = "0.1.1"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65"
|
||||
|
||||
[[package]]
|
||||
name = "vaultrs"
|
||||
version = "0.8.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "30ffcc0e81025065dda612ec1e26a3d81bb16ef3062354873d17a35965d68522"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"derive_builder",
|
||||
"http 1.4.2",
|
||||
"reqwest 0.13.5",
|
||||
"rustify",
|
||||
"rustify_derive",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"thiserror 2.0.19",
|
||||
"tracing",
|
||||
"url",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "vcpkg"
|
||||
version = "0.2.15"
|
||||
|
|
@ -5890,7 +6010,7 @@ dependencies = [
|
|||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
"synstructure",
|
||||
"synstructure 0.13.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5931,7 +6051,7 @@ dependencies = [
|
|||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
"synstructure",
|
||||
"synstructure 0.13.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ litellm-secrets = { path = "crates/secrets" }
|
|||
litellm-secrets-types = { path = "crates/secrets-types" }
|
||||
litellm-secrets-aws = { path = "crates/secrets-aws" }
|
||||
litellm-secrets-google = { path = "crates/secrets-google" }
|
||||
litellm-secrets-hashicorp = { path = "crates/secrets-hashicorp" }
|
||||
litellm-secrets-azure = { path = "crates/secrets-azure" }
|
||||
litellm-secrets-cyberark = { path = "crates/secrets-cyberark" }
|
||||
litellm-http = { path = "crates/http" }
|
||||
|
|
@ -32,6 +33,7 @@ litellm-cache = { path = "crates/cache" }
|
|||
litellm-cache-azure-blob = { path = "crates/cache-azure-blob" }
|
||||
litellm-cache-memory = { path = "crates/cache-memory" }
|
||||
litellm-cache-redis = { path = "crates/cache-redis" }
|
||||
litellm-cache-gcs = { path = "crates/cache-gcs" }
|
||||
litellm-cache-disk = { path = "crates/cache-disk" }
|
||||
litellm-cache-redis-semantic = { path = "crates/cache-redis-semantic" }
|
||||
litellm-cache-response = { path = "crates/cache-response" }
|
||||
|
|
@ -55,6 +57,9 @@ reqwest = { version = "0.12", default-features = false, features = ["json", "mul
|
|||
rstest = "0.26.1"
|
||||
rstest_reuse = "0.7.0"
|
||||
rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"] }
|
||||
rustify = "=0.7.0"
|
||||
rustify_derive = "=0.5.5"
|
||||
vaultrs = { version = "=0.8.0", default-features = false, features = ["rustls"] }
|
||||
rustls-native-certs = "0.8"
|
||||
serde = { version = "1.0", features = ["derive"] }
|
||||
serde_json = { version = "1.0", features = ["float_roundtrip"] }
|
||||
|
|
@ -71,6 +76,7 @@ base64 = "0.22"
|
|||
moka = { version = "0.12.16", features = ["future"] }
|
||||
strum = { version = "0.28.0", features = ["derive"] }
|
||||
url = "2.5.8"
|
||||
percent-encoding = "2.3"
|
||||
webpki-roots = "1"
|
||||
time = { version = "0.3.53", features = ["parsing"] }
|
||||
criterion = "0.8.2"
|
||||
|
|
|
|||
|
|
@ -128,6 +128,14 @@ impl VertexAuth {
|
|||
}
|
||||
}
|
||||
|
||||
pub async fn access_token(
|
||||
&self,
|
||||
config: &VertexConfig,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<String, Error> {
|
||||
self.load_provider(config, env_lookup).await?.token().await
|
||||
}
|
||||
|
||||
pub async fn validate_environment(
|
||||
&self,
|
||||
headers: Vec<(String, String)>,
|
||||
|
|
|
|||
20
litellm-rust/crates/cache-gcs/Cargo.toml
Normal file
20
litellm-rust/crates/cache-gcs/Cargo.toml
Normal file
|
|
@ -0,0 +1,20 @@
|
|||
[package]
|
||||
name = "litellm-cache-gcs"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
futures-util.workspace = true
|
||||
litellm-auth-gcp.workspace = true
|
||||
litellm-auth-types.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
percent-encoding.workspace = true
|
||||
reqwest.workspace = true
|
||||
tokio.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
260
litellm-rust/crates/cache-gcs/src/cache.rs
Normal file
260
litellm-rust/crates/cache-gcs/src/cache.rs
Normal file
|
|
@ -0,0 +1,260 @@
|
|||
use std::{future::Future, sync::Arc, time::Duration};
|
||||
|
||||
use futures_util::future::try_join_all;
|
||||
use litellm_cache::{
|
||||
BaseCache, BatchCache, BatchEntry, CacheCodec, CacheConnectionResult, Error, ExactCacheContext,
|
||||
FlushCache,
|
||||
};
|
||||
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, percent_encode};
|
||||
use reqwest::Client;
|
||||
|
||||
use crate::{GcpTokenSource, TokenSource};
|
||||
|
||||
pub const DEFAULT_ENDPOINT: &str = "https://storage.googleapis.com";
|
||||
|
||||
const OBJECT_NAME_ENCODE_SET: &AsciiSet = &NON_ALPHANUMERIC
|
||||
.remove(b'-')
|
||||
.remove(b'_')
|
||||
.remove(b'.')
|
||||
.remove(b'~');
|
||||
|
||||
pub fn key_prefix(gcs_path: Option<&str>) -> String {
|
||||
match gcs_path {
|
||||
Some(path) if !path.is_empty() => format!("{}/", path.trim_end_matches('/')),
|
||||
_ => String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct GcsConfig {
|
||||
pub bucket_name: String,
|
||||
pub gcs_path: Option<String>,
|
||||
pub path_service_account: Option<String>,
|
||||
pub endpoint: String,
|
||||
}
|
||||
|
||||
impl GcsConfig {
|
||||
pub fn new(bucket_name: impl Into<String>) -> Self {
|
||||
Self {
|
||||
bucket_name: bucket_name.into(),
|
||||
gcs_path: None,
|
||||
path_service_account: None,
|
||||
endpoint: DEFAULT_ENDPOINT.to_string(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct GcsCache<S: CacheCodec> {
|
||||
config: GcsConfig,
|
||||
key_prefix: String,
|
||||
client: Client,
|
||||
token: Arc<dyn TokenSource>,
|
||||
codec: S,
|
||||
}
|
||||
|
||||
impl<S: CacheCodec> GcsCache<S> {
|
||||
pub fn new(config: GcsConfig, codec: S) -> Result<Self, Error> {
|
||||
let token = Arc::new(GcpTokenSource::new(config.path_service_account.clone()));
|
||||
Self::with_token_source(config, codec, token)
|
||||
}
|
||||
|
||||
pub fn with_token_source(
|
||||
config: GcsConfig,
|
||||
codec: S,
|
||||
token: Arc<dyn TokenSource>,
|
||||
) -> Result<Self, Error> {
|
||||
let client = Client::builder().build().map_err(|_| Error::Unavailable)?;
|
||||
let key_prefix = key_prefix(config.gcs_path.as_deref());
|
||||
Ok(Self {
|
||||
config,
|
||||
key_prefix,
|
||||
client,
|
||||
token,
|
||||
codec,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn bucket_name(&self) -> &str {
|
||||
&self.config.bucket_name
|
||||
}
|
||||
|
||||
pub fn key_prefix(&self) -> &str {
|
||||
&self.key_prefix
|
||||
}
|
||||
|
||||
pub fn path_service_account(&self) -> Option<&str> {
|
||||
self.config.path_service_account.as_deref()
|
||||
}
|
||||
|
||||
pub fn object_name(&self, key: &str) -> String {
|
||||
format!("{}{}", self.key_prefix, key)
|
||||
}
|
||||
|
||||
fn encoded_object_name(&self, key: &str) -> String {
|
||||
percent_encode(self.object_name(key).as_bytes(), OBJECT_NAME_ENCODE_SET).to_string()
|
||||
}
|
||||
|
||||
fn endpoint(&self, path: &str) -> String {
|
||||
format!("{}{}", self.config.endpoint.trim_end_matches('/'), path)
|
||||
}
|
||||
|
||||
async fn async_set(&self, key: &str, value: S::Value) -> Result<(), Error> {
|
||||
let token = self.token.bearer_token().await?;
|
||||
let payload = self.codec.encode(&value)?;
|
||||
let url = self.endpoint(&format!(
|
||||
"/upload/storage/v1/b/{}/o?uploadType=media&name={}",
|
||||
self.config.bucket_name,
|
||||
self.encoded_object_name(key)
|
||||
));
|
||||
let response = self
|
||||
.client
|
||||
.post(url)
|
||||
.bearer_auth(token)
|
||||
.header(reqwest::header::CONTENT_TYPE, "application/json")
|
||||
.body(payload)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::Unavailable);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn async_get(&self, key: &str) -> Result<Option<S::Value>, Error> {
|
||||
let token = self.token.bearer_token().await?;
|
||||
let url = self.endpoint(&format!(
|
||||
"/storage/v1/b/{}/o/{}?alt=media",
|
||||
self.config.bucket_name,
|
||||
self.encoded_object_name(key)
|
||||
));
|
||||
let response = self
|
||||
.client
|
||||
.get(url)
|
||||
.bearer_auth(token)
|
||||
.header(reqwest::header::CONTENT_TYPE, "application/json")
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
if response.status() == reqwest::StatusCode::NOT_FOUND {
|
||||
return Ok(None);
|
||||
}
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::Unavailable);
|
||||
}
|
||||
let body = response.bytes().await.map_err(|_| Error::Unavailable)?;
|
||||
self.codec
|
||||
.decode(&body)
|
||||
.map(Some)
|
||||
.map_err(|_| Error::InvalidEntry)
|
||||
}
|
||||
|
||||
fn run_sync<T, F>(future: F) -> Result<T, Error>
|
||||
where
|
||||
F: Future<Output = Result<T, Error>> + Send,
|
||||
T: Send,
|
||||
{
|
||||
let run = || {
|
||||
tokio::runtime::Builder::new_current_thread()
|
||||
.enable_all()
|
||||
.build()
|
||||
.map_err(|_| Error::Unavailable)
|
||||
.and_then(|runtime| runtime.block_on(future))
|
||||
};
|
||||
if let Ok(handle) = tokio::runtime::Handle::try_current() {
|
||||
if handle.runtime_flavor() == tokio::runtime::RuntimeFlavor::MultiThread {
|
||||
return tokio::task::block_in_place(run);
|
||||
}
|
||||
return std::thread::scope(|scope| {
|
||||
scope
|
||||
.spawn(run)
|
||||
.join()
|
||||
.map_err(|_| Error::Unavailable)
|
||||
.and_then(|result| result)
|
||||
});
|
||||
}
|
||||
run()
|
||||
}
|
||||
}
|
||||
|
||||
impl<S: CacheCodec> BaseCache for GcsCache<S> {
|
||||
type Value = S::Value;
|
||||
type Context = ExactCacheContext;
|
||||
|
||||
fn get_ttl(&self, _: &Self::Context) -> Option<Duration> {
|
||||
None
|
||||
}
|
||||
|
||||
fn set_cache(&self, key: &str, value: Self::Value, _: &Self::Context) -> Result<(), Error> {
|
||||
Self::run_sync(self.async_set(key, value))
|
||||
}
|
||||
|
||||
fn get_cache(&self, key: &str, _: &Self::Context) -> Result<Option<Self::Value>, Error> {
|
||||
Self::run_sync(self.async_get(key))
|
||||
}
|
||||
|
||||
async fn async_set_cache(
|
||||
&self,
|
||||
key: &str,
|
||||
value: Self::Value,
|
||||
_: Self::Context,
|
||||
) -> Result<(), Error> {
|
||||
self.async_set(key, value).await
|
||||
}
|
||||
|
||||
async fn async_get_cache(
|
||||
&self,
|
||||
key: &str,
|
||||
_: &Self::Context,
|
||||
) -> Result<Option<Self::Value>, Error> {
|
||||
self.async_get(key).await
|
||||
}
|
||||
|
||||
async fn async_set_cache_pipeline(
|
||||
&self,
|
||||
entries: Vec<(String, Self::Value)>,
|
||||
context: Self::Context,
|
||||
) -> Result<(), Error> {
|
||||
try_join_all(entries.into_iter().map(|(key, value)| {
|
||||
let context = context.clone();
|
||||
async move { self.async_set_cache(&key, value, context).await }
|
||||
}))
|
||||
.await
|
||||
.map(|_| ())
|
||||
}
|
||||
|
||||
async fn disconnect(&self) -> Result<(), Error> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
|
||||
Err(Error::UnsupportedOperation)
|
||||
}
|
||||
}
|
||||
|
||||
impl<S: CacheCodec> BatchCache for GcsCache<S> {
|
||||
async fn async_batch_get_cache(
|
||||
&self,
|
||||
keys: Vec<String>,
|
||||
context: Self::Context,
|
||||
) -> Result<Vec<BatchEntry<Self::Value>>, Error> {
|
||||
try_join_all(keys.into_iter().map(|key| {
|
||||
let context = context.clone();
|
||||
async move {
|
||||
match self.async_get_cache(&key, &context).await {
|
||||
Ok(Some(value)) => Ok(BatchEntry::Hit(value)),
|
||||
Ok(None) => Ok(BatchEntry::Miss),
|
||||
Err(Error::InvalidEntry) => Ok(BatchEntry::Invalid),
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
}
|
||||
}))
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
impl<S: CacheCodec> FlushCache for GcsCache<S> {
|
||||
fn flush_cache(&self) -> Result<(), Error> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
5
litellm-rust/crates/cache-gcs/src/lib.rs
Normal file
5
litellm-rust/crates/cache-gcs/src/lib.rs
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
mod cache;
|
||||
mod token;
|
||||
|
||||
pub use cache::{DEFAULT_ENDPOINT, GcsCache, GcsConfig, key_prefix};
|
||||
pub use token::{GcpTokenSource, StaticTokenSource, TokenSource};
|
||||
44
litellm-rust/crates/cache-gcs/src/token.rs
Normal file
44
litellm-rust/crates/cache-gcs/src/token.rs
Normal file
|
|
@ -0,0 +1,44 @@
|
|||
use std::{future::Future, pin::Pin};
|
||||
|
||||
use litellm_auth_gcp::{VertexAuth, VertexConfig};
|
||||
use litellm_auth_types::{InputSource, SecretValue, Sourced};
|
||||
use litellm_cache::Error;
|
||||
|
||||
pub trait TokenSource: Send + Sync + 'static {
|
||||
fn bearer_token(&self) -> Pin<Box<dyn Future<Output = Result<String, Error>> + Send + '_>>;
|
||||
}
|
||||
|
||||
pub struct GcpTokenSource {
|
||||
auth: VertexAuth,
|
||||
config: VertexConfig,
|
||||
}
|
||||
|
||||
impl GcpTokenSource {
|
||||
pub fn new(path_service_account: Option<String>) -> Self {
|
||||
let credentials = path_service_account
|
||||
.map(|path| Sourced::new(SecretValue::new(path), InputSource::Deployment));
|
||||
Self {
|
||||
auth: VertexAuth::default(),
|
||||
config: VertexConfig::new(credentials, None, None),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl TokenSource for GcpTokenSource {
|
||||
fn bearer_token(&self) -> Pin<Box<dyn Future<Output = Result<String, Error>> + Send + '_>> {
|
||||
Box::pin(async move {
|
||||
self.auth
|
||||
.access_token(&self.config, &|name| std::env::var(name).ok())
|
||||
.await
|
||||
.map_err(|_| Error::Unavailable)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub struct StaticTokenSource(pub String);
|
||||
|
||||
impl TokenSource for StaticTokenSource {
|
||||
fn bearer_token(&self) -> Pin<Box<dyn Future<Output = Result<String, Error>> + Send + '_>> {
|
||||
Box::pin(async move { Ok(self.0.clone()) })
|
||||
}
|
||||
}
|
||||
324
litellm-rust/crates/cache-gcs/tests/cache.rs
Normal file
324
litellm-rust/crates/cache-gcs/tests/cache.rs
Normal file
|
|
@ -0,0 +1,324 @@
|
|||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use litellm_cache::{
|
||||
BaseCache, BatchCache, BatchEntry, CacheContext, Error, ExactCacheContext, FlushCache,
|
||||
JsonCodec,
|
||||
};
|
||||
use litellm_cache_gcs::{GcsCache, GcsConfig, StaticTokenSource, TokenSource, key_prefix};
|
||||
use serde_json::json;
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{body_bytes, header, method, path, query_param},
|
||||
};
|
||||
|
||||
fn config(server: &MockServer, gcs_path: Option<&str>) -> GcsConfig {
|
||||
GcsConfig {
|
||||
bucket_name: "bucket".into(),
|
||||
gcs_path: gcs_path.map(str::to_string),
|
||||
path_service_account: None,
|
||||
endpoint: server.uri(),
|
||||
}
|
||||
}
|
||||
|
||||
fn cache(server: &MockServer, gcs_path: Option<&str>) -> GcsCache<JsonCodec<serde_json::Value>> {
|
||||
GcsCache::with_token_source(
|
||||
config(server, gcs_path),
|
||||
JsonCodec::new(),
|
||||
Arc::new(StaticTokenSource("tok".into())),
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn set_writes_encoded_object_and_headers() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/upload/storage/v1/b/bucket/o"))
|
||||
.and(query_param("uploadType", "media"))
|
||||
.and(header("authorization", "Bearer tok"))
|
||||
.and(header("content-type", "application/json"))
|
||||
.and(body_bytes(br#"{"value":"entry"}"#))
|
||||
.respond_with(ResponseTemplate::new(200))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
cache(&server, Some("cache/"))
|
||||
.set_cache(
|
||||
"team:a b/c",
|
||||
json!({"value": "entry"}),
|
||||
&ExactCacheContext::default(),
|
||||
)
|
||||
.unwrap();
|
||||
let requests = server.received_requests().await.unwrap();
|
||||
assert_eq!(requests.len(), 1);
|
||||
assert_eq!(
|
||||
requests[0].url.query(),
|
||||
Some("uploadType=media&name=cache%2Fteam%3Aa%20b%2Fc")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_maps_statuses_and_decode_failures() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/hit"))
|
||||
.and(query_param("alt", "media"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"value": "entry"})))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/missing"))
|
||||
.respond_with(ResponseTemplate::new(404))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/server-error"))
|
||||
.respond_with(ResponseTemplate::new(500))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/invalid"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_string("not json"))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let cache = cache(&server, None);
|
||||
assert_eq!(
|
||||
cache
|
||||
.get_cache("hit", &ExactCacheContext::default())
|
||||
.unwrap(),
|
||||
Some(json!({"value": "entry"}))
|
||||
);
|
||||
assert_eq!(
|
||||
cache
|
||||
.get_cache("missing", &ExactCacheContext::default())
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
cache
|
||||
.get_cache("server-error", &ExactCacheContext::default())
|
||||
.unwrap_err(),
|
||||
Error::Unavailable
|
||||
);
|
||||
assert_eq!(
|
||||
cache
|
||||
.get_cache("invalid", &ExactCacheContext::default())
|
||||
.unwrap_err(),
|
||||
Error::InvalidEntry
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn key_prefix_normalizes_paths() {
|
||||
assert_eq!(key_prefix(None), "");
|
||||
assert_eq!(key_prefix(Some("a/b/")), "a/b/");
|
||||
assert_eq!(key_prefix(Some("a/b")), "a/b/");
|
||||
assert_eq!(key_prefix(Some("")), "");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn object_names_use_python_quote_encoding() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/upload/storage/v1/b/bucket/o"))
|
||||
.and(query_param("uploadType", "media"))
|
||||
.respond_with(ResponseTemplate::new(200))
|
||||
.expect(2)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let cache = cache(&server, Some("p/"));
|
||||
cache
|
||||
.set_cache(
|
||||
"a~b-c_d.e/f g%h",
|
||||
json!({"value": "punctuation"}),
|
||||
&ExactCacheContext::default(),
|
||||
)
|
||||
.unwrap();
|
||||
cache
|
||||
.set_cache(
|
||||
"ключ",
|
||||
json!({"value": "utf8"}),
|
||||
&ExactCacheContext::default(),
|
||||
)
|
||||
.unwrap();
|
||||
let requests = server.received_requests().await.unwrap();
|
||||
let queries: Vec<_> = requests
|
||||
.iter()
|
||||
.filter_map(|request| request.url.query())
|
||||
.collect();
|
||||
assert!(queries.contains(&"uploadType=media&name=p%2Fa~b-c_d.e%2Ff%20g%25h"));
|
||||
assert!(queries.contains(&"uploadType=media&name=p%2F%D0%BA%D0%BB%D1%8E%D1%87"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ignores_ttl_and_writes_pipeline_concurrently() {
|
||||
let server = MockServer::start().await;
|
||||
for key in ["one", "two", "three"] {
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/upload/storage/v1/b/bucket/o"))
|
||||
.and(query_param("uploadType", "media"))
|
||||
.and(query_param("name", key))
|
||||
.respond_with(ResponseTemplate::new(200))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
}
|
||||
let cache = cache(&server, None);
|
||||
assert_eq!(cache.get_ttl(&ExactCacheContext::default()), None);
|
||||
assert_eq!(
|
||||
cache.get_ttl(&ExactCacheContext::default().with_ttl(Some(Duration::from_secs(5)))),
|
||||
None
|
||||
);
|
||||
cache
|
||||
.async_set_cache_pipeline(
|
||||
vec![
|
||||
("one".into(), json!({"key": "one"})),
|
||||
("two".into(), json!({"key": "two"})),
|
||||
("three".into(), json!({"key": "three"})),
|
||||
],
|
||||
ExactCacheContext::default().with_ttl(Some(Duration::from_secs(5))),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn async_batch_get_preserves_hits_misses_and_invalid_entries() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/hit"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"value": "entry"})))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/missing"))
|
||||
.respond_with(ResponseTemplate::new(404))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/invalid"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_string("not json"))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
assert_eq!(
|
||||
cache(&server, None)
|
||||
.async_batch_get_cache(
|
||||
vec!["hit".into(), "missing".into(), "invalid".into()],
|
||||
ExactCacheContext::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
vec![
|
||||
BatchEntry::Hit(json!({"value": "entry"})),
|
||||
BatchEntry::Miss,
|
||||
BatchEntry::Invalid,
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn lifecycle_operations_are_noops_and_connection_test_is_unsupported() {
|
||||
let server = MockServer::start().await;
|
||||
let cache = cache(&server, None);
|
||||
assert_eq!(cache.flush_cache(), Ok(()));
|
||||
assert_eq!(cache.disconnect().await, Ok(()));
|
||||
assert_eq!(
|
||||
cache.test_connection().await,
|
||||
Err(Error::UnsupportedOperation)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sync_operations_work_without_an_active_runtime() {
|
||||
let runtime = tokio::runtime::Builder::new_multi_thread()
|
||||
.enable_all()
|
||||
.build()
|
||||
.unwrap();
|
||||
let server = runtime.block_on(MockServer::start());
|
||||
runtime.block_on(
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/upload/storage/v1/b/bucket/o"))
|
||||
.respond_with(ResponseTemplate::new(200))
|
||||
.mount(&server),
|
||||
);
|
||||
runtime.block_on(
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/key"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"value": "entry"})))
|
||||
.mount(&server),
|
||||
);
|
||||
let cache = cache(&server, None);
|
||||
cache
|
||||
.set_cache(
|
||||
"key",
|
||||
json!({"value": "entry"}),
|
||||
&ExactCacheContext::default(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
cache
|
||||
.get_cache("key", &ExactCacheContext::default())
|
||||
.unwrap(),
|
||||
Some(json!({"value": "entry"}))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn sync_operations_work_inside_a_multi_thread_runtime() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/upload/storage/v1/b/bucket/o"))
|
||||
.respond_with(ResponseTemplate::new(200))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/key"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"value": "entry"})))
|
||||
.mount(&server)
|
||||
.await;
|
||||
let cache = cache(&server, None);
|
||||
cache
|
||||
.set_cache(
|
||||
"key",
|
||||
json!({"value": "entry"}),
|
||||
&ExactCacheContext::default(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
cache
|
||||
.get_cache("key", &ExactCacheContext::default())
|
||||
.unwrap(),
|
||||
Some(json!({"value": "entry"}))
|
||||
);
|
||||
}
|
||||
|
||||
struct FailingTokenSource;
|
||||
|
||||
impl TokenSource for FailingTokenSource {
|
||||
fn bearer_token(
|
||||
&self,
|
||||
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<String, Error>> + Send + '_>>
|
||||
{
|
||||
Box::pin(async { Err(Error::Unavailable) })
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn token_source_failure_skips_http() {
|
||||
let server = MockServer::start().await;
|
||||
let cache = GcsCache::with_token_source(
|
||||
config(&server, None),
|
||||
JsonCodec::<serde_json::Value>::new(),
|
||||
Arc::new(FailingTokenSource),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
cache
|
||||
.get_cache("key", &ExactCacheContext::default())
|
||||
.unwrap_err(),
|
||||
Error::Unavailable
|
||||
);
|
||||
assert_eq!(server.received_requests().await.unwrap().len(), 0);
|
||||
}
|
||||
2
litellm-rust/crates/cache/src/error.rs
vendored
2
litellm-rust/crates/cache/src/error.rs
vendored
|
|
@ -6,6 +6,6 @@ pub enum Error {
|
|||
InvalidEntry,
|
||||
#[error("flushing Redis requires an explicit namespace")]
|
||||
UnscopedFlush,
|
||||
#[error("cache backend does not support this operation")]
|
||||
#[error("operation is not supported by this cache")]
|
||||
UnsupportedOperation,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ litellm-cache.workspace = true
|
|||
litellm-cache-azure-blob.workspace = true
|
||||
litellm-cache-memory.workspace = true
|
||||
litellm-cache-redis.workspace = true
|
||||
litellm-cache-gcs.workspace = true
|
||||
litellm-cache-disk.workspace = true
|
||||
litellm-cache-redis-semantic.workspace = true
|
||||
litellm-cache-response.workspace = true
|
||||
|
|
|
|||
|
|
@ -79,6 +79,13 @@ pub(super) struct RedisCacheConfig {
|
|||
pub(super) connection: RedisConnectionConfig,
|
||||
}
|
||||
|
||||
#[derive(Debug, PartialEq)]
|
||||
pub(super) struct GcsCacheConfig {
|
||||
pub(super) bucket_name: String,
|
||||
pub(super) key_prefix: String,
|
||||
pub(super) path_service_account: Option<String>,
|
||||
}
|
||||
|
||||
pub(super) struct AzureBlobCacheConfig {
|
||||
pub(super) account_url: String,
|
||||
pub(super) container: String,
|
||||
|
|
@ -111,6 +118,7 @@ const REDIS_PY_DEFAULT_MAX_CONNECTIONS: usize = 1 << 31;
|
|||
pub(super) enum CacheBackendConfig {
|
||||
Memory(MemoryCacheConfig),
|
||||
Redis(Box<RedisCacheConfig>),
|
||||
Gcs(GcsCacheConfig),
|
||||
Disk(DiskCacheConfig),
|
||||
AzureBlob(AzureBlobCacheConfig),
|
||||
RedisSemantic(Box<RedisSemanticCacheConfig>),
|
||||
|
|
@ -128,6 +136,7 @@ pub(super) enum UnsupportedCacheConfig {
|
|||
RedisCredentials,
|
||||
RedisConnection,
|
||||
RedisOption,
|
||||
GcsBucket,
|
||||
DiskStore,
|
||||
}
|
||||
|
||||
|
|
@ -139,6 +148,7 @@ impl UnsupportedCacheConfig {
|
|||
Self::RedisCredentials => "native Redis credentials require Python",
|
||||
Self::RedisConnection => "native Redis connection type is not implemented",
|
||||
Self::RedisOption => "native Redis configuration requires Python",
|
||||
Self::GcsBucket => "native GCS cache requires a configured bucket name",
|
||||
Self::DiskStore => "native disk cache requires the built-in diskcache store",
|
||||
}
|
||||
}
|
||||
|
|
@ -182,6 +192,13 @@ impl NativeCacheConfig {
|
|||
}))),
|
||||
Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)),
|
||||
},
|
||||
Some(CacheType::Gcs) => match project_gcs(&backend)? {
|
||||
Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self {
|
||||
policy,
|
||||
backend: CacheBackendConfig::Gcs(backend),
|
||||
}))),
|
||||
Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)),
|
||||
},
|
||||
Some(CacheType::Disk) => match project_disk(&backend)? {
|
||||
Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self {
|
||||
policy,
|
||||
|
|
@ -201,15 +218,11 @@ impl NativeCacheConfig {
|
|||
backend: CacheBackendConfig::RedisSemantic(Box::new(backend)),
|
||||
}))
|
||||
}),
|
||||
Some(
|
||||
CacheType::ValkeySemantic
|
||||
| CacheType::S3
|
||||
| CacheType::QdrantSemantic
|
||||
| CacheType::Gcs,
|
||||
)
|
||||
| None => Ok(CacheConfigProjection::Unsupported(
|
||||
UnsupportedCacheConfig::Backend,
|
||||
)),
|
||||
Some(CacheType::ValkeySemantic | CacheType::S3 | CacheType::QdrantSemantic) | None => {
|
||||
Ok(CacheConfigProjection::Unsupported(
|
||||
UnsupportedCacheConfig::Backend,
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -217,8 +230,10 @@ impl NativeCacheConfig {
|
|||
let default_ttl = match &self.backend {
|
||||
CacheBackendConfig::Memory(config) => Some(config.default_ttl),
|
||||
CacheBackendConfig::Redis(config) => Some(config.default_ttl),
|
||||
CacheBackendConfig::Disk(_) => None,
|
||||
CacheBackendConfig::AzureBlob(_) | CacheBackendConfig::RedisSemantic(_) => None,
|
||||
CacheBackendConfig::Disk(_)
|
||||
| CacheBackendConfig::AzureBlob(_)
|
||||
| CacheBackendConfig::Gcs(_)
|
||||
| CacheBackendConfig::RedisSemantic(_) => None,
|
||||
};
|
||||
if service.default_ttl() != default_ttl {
|
||||
return Some("facade and native backend default TTLs must match");
|
||||
|
|
@ -245,6 +260,31 @@ impl NativeCacheConfig {
|
|||
CacheBackendConfig::Redis(config) => (service.namespace()
|
||||
!= config.namespace.as_deref())
|
||||
.then_some("facade and native backend namespaces must match"),
|
||||
CacheBackendConfig::Gcs(_) if service.kind() != "gcs" => {
|
||||
Some("facade and native backend types must match")
|
||||
}
|
||||
CacheBackendConfig::Gcs(config)
|
||||
if service
|
||||
.gcs_backend()
|
||||
.is_none_or(|backend| backend.bucket_name() != config.bucket_name) =>
|
||||
{
|
||||
Some("facade and native backend buckets must match")
|
||||
}
|
||||
CacheBackendConfig::Gcs(config)
|
||||
if service
|
||||
.gcs_backend()
|
||||
.is_none_or(|backend| backend.key_prefix() != config.key_prefix) =>
|
||||
{
|
||||
Some("facade and native backend key prefixes must match")
|
||||
}
|
||||
CacheBackendConfig::Gcs(config)
|
||||
if service.gcs_backend().is_none_or(|backend| {
|
||||
backend.path_service_account() != config.path_service_account.as_deref()
|
||||
}) =>
|
||||
{
|
||||
Some("facade and native backend credentials must match")
|
||||
}
|
||||
CacheBackendConfig::Gcs(_) => None,
|
||||
CacheBackendConfig::Disk(_) if service.kind() != "disk" => {
|
||||
Some("facade and native backend types must match")
|
||||
}
|
||||
|
|
@ -331,6 +371,23 @@ fn project_memory(backend: &Bound<'_, PyAny>) -> PyResult<MemoryCacheConfig> {
|
|||
})
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
fn project_gcs(
|
||||
backend: &Bound<'_, PyAny>,
|
||||
) -> PyResult<Result<GcsCacheConfig, UnsupportedCacheConfig>> {
|
||||
let bucket_name = match backend.getattr("bucket_name")?.extract::<Option<String>>() {
|
||||
Ok(Some(bucket_name)) if !bucket_name.is_empty() => bucket_name,
|
||||
_ => return Ok(Err(UnsupportedCacheConfig::GcsBucket)),
|
||||
};
|
||||
Ok(Ok(GcsCacheConfig {
|
||||
bucket_name,
|
||||
key_prefix: backend.getattr("key_prefix")?.extract::<String>()?,
|
||||
path_service_account: backend
|
||||
.getattr("path_service_account")?
|
||||
.extract::<Option<String>>()?,
|
||||
}))
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
fn project_disk(
|
||||
backend: &Bound<'_, PyAny>,
|
||||
|
|
@ -734,7 +791,7 @@ mod tests {
|
|||
|
||||
use super::{
|
||||
CacheBackendConfig, CacheConfigProjection, CachePolicy, CertificateRequirement,
|
||||
DiskCacheConfig, NativeCacheConfig, RedisProtocol,
|
||||
DiskCacheConfig, GcsCacheConfig, NativeCacheConfig, RedisProtocol, UnsupportedCacheConfig,
|
||||
};
|
||||
use crate::cache::native::NativeResponseCache;
|
||||
|
||||
|
|
@ -809,6 +866,71 @@ mod tests {
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn projects_gcs_configuration() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let facade = facade(
|
||||
py,
|
||||
"backend = SimpleNamespace(bucket_name='bucket', key_prefix='cache/', path_service_account='credentials.json')\n\
|
||||
facade = SimpleNamespace(type='gcs', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)",
|
||||
);
|
||||
let CacheConfigProjection::Native(config) =
|
||||
NativeCacheConfig::project(&facade).unwrap()
|
||||
else {
|
||||
panic!("GCS cache should be supported");
|
||||
};
|
||||
let CacheBackendConfig::Gcs(gcs) = config.backend else {
|
||||
panic!("expected GCS configuration");
|
||||
};
|
||||
assert_eq!(
|
||||
gcs,
|
||||
GcsCacheConfig {
|
||||
bucket_name: "bucket".into(),
|
||||
key_prefix: "cache/".into(),
|
||||
path_service_account: Some("credentials.json".into()),
|
||||
}
|
||||
);
|
||||
let matching = NativeResponseCache::gcs(
|
||||
litellm_cache_gcs::GcsConfig {
|
||||
bucket_name: "bucket".into(),
|
||||
gcs_path: Some("cache/".into()),
|
||||
path_service_account: Some("credentials.json".into()),
|
||||
endpoint: litellm_cache_gcs::DEFAULT_ENDPOINT.into(),
|
||||
},
|
||||
Some("token".into()),
|
||||
)
|
||||
.unwrap();
|
||||
let matching_config = NativeCacheConfig {
|
||||
policy: config.policy,
|
||||
backend: CacheBackendConfig::Gcs(gcs),
|
||||
};
|
||||
assert_eq!(matching_config.service_mismatch(&matching), None);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_gcs_without_a_bucket_name() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let facade = facade(
|
||||
py,
|
||||
"backend = SimpleNamespace(bucket_name=None, key_prefix='', path_service_account=None)\n\
|
||||
facade = SimpleNamespace(type='gcs', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)",
|
||||
);
|
||||
let CacheConfigProjection::Unsupported(reason) =
|
||||
NativeCacheConfig::project(&facade).unwrap()
|
||||
else {
|
||||
panic!("GCS cache without a bucket should be unsupported");
|
||||
};
|
||||
assert!(matches!(&reason, UnsupportedCacheConfig::GcsBucket));
|
||||
assert_eq!(
|
||||
reason.message(),
|
||||
"native GCS cache requires a configured bucket name"
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn projects_resolved_redis_tls_configuration() {
|
||||
Python::initialize();
|
||||
|
|
|
|||
|
|
@ -328,6 +328,7 @@ impl FacadeGuard {
|
|||
"RedisClusterCache",
|
||||
"redis",
|
||||
),
|
||||
("gcs", _) => ("litellm.caching.gcs_cache", "GCSCache", "gcs"),
|
||||
("disk", _) => ("litellm.caching.disk_cache", "DiskCache", "disk"),
|
||||
("azure-blob", _) => (
|
||||
"litellm.caching.azure_blob_cache",
|
||||
|
|
@ -393,6 +394,9 @@ impl FacadeGuard {
|
|||
"embedding_timeout",
|
||||
"_index_name",
|
||||
"_redis_url",
|
||||
"bucket_name",
|
||||
"key_prefix",
|
||||
"path_service_account",
|
||||
],
|
||||
)?,
|
||||
disk_store: (kind == "disk")
|
||||
|
|
|
|||
|
|
@ -7,6 +7,8 @@ use pyo3::{
|
|||
prelude::*,
|
||||
};
|
||||
|
||||
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
|
||||
|
||||
use super::{
|
||||
cache_error, config::project_redis_semantic, embedder::PythonEmbedder, facade::FacadeGuard,
|
||||
native::NativeResponseCache, request::duration,
|
||||
|
|
@ -72,6 +74,31 @@ impl CacheTestHandle {
|
|||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (bucket_name, *, gcs_path=None, path_service_account=None, endpoint=None, token=None))]
|
||||
fn gcs(
|
||||
py: Python<'_>,
|
||||
bucket_name: String,
|
||||
gcs_path: Option<String>,
|
||||
path_service_account: Option<String>,
|
||||
endpoint: Option<String>,
|
||||
token: Option<String>,
|
||||
) -> PyResult<Self> {
|
||||
let config = GcsConfig {
|
||||
bucket_name,
|
||||
gcs_path,
|
||||
path_service_account,
|
||||
endpoint: endpoint.unwrap_or_else(|| DEFAULT_ENDPOINT.to_string()),
|
||||
};
|
||||
let service = release_gil(py, move || NativeResponseCache::gcs(config, token))
|
||||
.map_err(cache_error)?;
|
||||
Ok(Self {
|
||||
service,
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (directory))]
|
||||
fn disk(py: Python<'_>, directory: String) -> PyResult<Self> {
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ use std::{path::Path, sync::Arc, time::Duration};
|
|||
use litellm_cache::{CacheCodec, CacheConnectionResult, Error, ExactCacheContext};
|
||||
use litellm_cache_azure_blob::AzureBlobCache;
|
||||
use litellm_cache_disk::DiskCache;
|
||||
use litellm_cache_gcs::{GcsCache, GcsConfig, StaticTokenSource};
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_redis::{RedisCache, RedisTopology};
|
||||
use litellm_cache_redis_semantic::{RedisSemanticCache, RedisSemanticConfig};
|
||||
|
|
@ -21,6 +22,7 @@ pub(super) enum NativeResponseCache {
|
|||
cache: Arc<ResponseCache<RedisCache<ResponseCacheCodec>>>,
|
||||
buffer: Option<Arc<WriteBuffer>>,
|
||||
},
|
||||
Gcs(Arc<ResponseCache<GcsCache<ResponseCacheCodec>>>),
|
||||
Disk(Arc<ResponseCache<DiskCache<ResponseCacheCodec>>>),
|
||||
AzureBlob(Arc<ResponseCache<AzureBlobCache<ResponseCacheCodec>>>),
|
||||
RedisSemantic(Arc<ResponseCache<RedisSemanticCache<PythonEmbedder>>>),
|
||||
|
|
@ -70,6 +72,18 @@ impl NativeResponseCache {
|
|||
)))))
|
||||
}
|
||||
|
||||
pub fn gcs(config: GcsConfig, token: Option<String>) -> Result<Self, Error> {
|
||||
let backend = match token {
|
||||
Some(token) => GcsCache::with_token_source(
|
||||
config,
|
||||
ResponseCacheCodec,
|
||||
Arc::new(StaticTokenSource(token)),
|
||||
)?,
|
||||
None => GcsCache::new(config, ResponseCacheCodec)?,
|
||||
};
|
||||
Ok(Self::Gcs(Arc::new(ResponseCache::new(Arc::new(backend)))))
|
||||
}
|
||||
|
||||
pub async fn azure_blob(account_url: &str, container: &str) -> Result<Self, Error> {
|
||||
let backend = AzureBlobCache::connect(
|
||||
account_url,
|
||||
|
|
@ -89,7 +103,11 @@ impl NativeResponseCache {
|
|||
cache.backend().account_url(),
|
||||
cache.backend().container_name(),
|
||||
)),
|
||||
Self::Memory(_) | Self::Redis { .. } | Self::Disk(_) | Self::RedisSemantic(_) => None,
|
||||
Self::Memory(_)
|
||||
| Self::Redis { .. }
|
||||
| Self::Disk(_)
|
||||
| Self::Gcs(_)
|
||||
| Self::RedisSemantic(_) => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -99,6 +117,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Memory(_) => "memory",
|
||||
Self::Redis { .. } => "redis",
|
||||
Self::Gcs(_) => "gcs",
|
||||
Self::Disk(_) => "disk",
|
||||
Self::RedisSemantic(_) => "redis_semantic",
|
||||
Self::AzureBlob(_) => "azure-blob",
|
||||
|
|
@ -109,6 +128,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Memory(cache) => cache.default_ttl(),
|
||||
Self::Redis { cache, .. } => cache.default_ttl(),
|
||||
Self::Gcs(cache) => cache.default_ttl(),
|
||||
Self::Disk(cache) => cache.default_ttl(),
|
||||
Self::RedisSemantic(cache) => cache.default_ttl(),
|
||||
Self::AzureBlob(cache) => cache.default_ttl(),
|
||||
|
|
@ -119,12 +139,17 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Memory(_) | Self::Disk(_) | Self::AzureBlob(_) | Self::RedisSemantic(_) => None,
|
||||
Self::Redis { cache, .. } => cache.backend().namespace(),
|
||||
Self::Gcs(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn topology(&self) -> Option<&RedisTopology> {
|
||||
match self {
|
||||
Self::Memory(_) | Self::Disk(_) | Self::AzureBlob(_) | Self::RedisSemantic(_) => None,
|
||||
Self::Memory(_)
|
||||
| Self::Disk(_)
|
||||
| Self::AzureBlob(_)
|
||||
| Self::Gcs(_)
|
||||
| Self::RedisSemantic(_) => None,
|
||||
Self::Redis { cache, .. } => Some(cache.backend().topology()),
|
||||
}
|
||||
}
|
||||
|
|
@ -132,46 +157,66 @@ impl NativeResponseCache {
|
|||
pub fn capacity(&self) -> Option<usize> {
|
||||
match self {
|
||||
Self::Memory(cache) => Some(cache.backend().max_size_in_memory()),
|
||||
Self::Redis { .. } | Self::Disk(_) | Self::AzureBlob(_) | Self::RedisSemantic(_) => {
|
||||
None
|
||||
}
|
||||
Self::Redis { .. }
|
||||
| Self::Disk(_)
|
||||
| Self::AzureBlob(_)
|
||||
| Self::Gcs(_)
|
||||
| Self::RedisSemantic(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn max_entry_bytes(&self) -> Option<usize> {
|
||||
match self {
|
||||
Self::Memory(cache) => cache.backend().max_entry_bytes(),
|
||||
Self::Redis { .. } | Self::Disk(_) | Self::AzureBlob(_) | Self::RedisSemantic(_) => {
|
||||
None
|
||||
}
|
||||
Self::Redis { .. }
|
||||
| Self::Disk(_)
|
||||
| Self::AzureBlob(_)
|
||||
| Self::Gcs(_)
|
||||
| Self::RedisSemantic(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn index_name(&self) -> Option<&str> {
|
||||
match self {
|
||||
Self::RedisSemantic(cache) => Some(cache.backend().index_name()),
|
||||
Self::Memory(_) | Self::Redis { .. } | Self::AzureBlob(_) | Self::Disk(_) => None,
|
||||
Self::Memory(_)
|
||||
| Self::Redis { .. }
|
||||
| Self::AzureBlob(_)
|
||||
| Self::Disk(_)
|
||||
| Self::Gcs(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn similarity_threshold(&self) -> Option<f32> {
|
||||
match self {
|
||||
Self::RedisSemantic(cache) => Some(cache.backend().similarity_threshold()),
|
||||
Self::Memory(_) | Self::Redis { .. } | Self::AzureBlob(_) | Self::Disk(_) => None,
|
||||
Self::Memory(_)
|
||||
| Self::Redis { .. }
|
||||
| Self::AzureBlob(_)
|
||||
| Self::Disk(_)
|
||||
| Self::Gcs(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn semantic_embedder(&self) -> Option<&PythonEmbedder> {
|
||||
match self {
|
||||
Self::RedisSemantic(cache) => Some(cache.backend().embedder()),
|
||||
Self::Memory(_) | Self::Redis { .. } | Self::AzureBlob(_) | Self::Disk(_) => None,
|
||||
Self::Memory(_)
|
||||
| Self::Redis { .. }
|
||||
| Self::AzureBlob(_)
|
||||
| Self::Disk(_)
|
||||
| Self::Gcs(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn embedder_object(&self) -> Option<&Py<PyAny>> {
|
||||
match self {
|
||||
Self::RedisSemantic(cache) => Some(cache.backend().embedder().object()),
|
||||
Self::Memory(_) | Self::Redis { .. } | Self::AzureBlob(_) | Self::Disk(_) => None,
|
||||
Self::Memory(_)
|
||||
| Self::Redis { .. }
|
||||
| Self::AzureBlob(_)
|
||||
| Self::Disk(_)
|
||||
| Self::Gcs(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -195,9 +240,11 @@ impl NativeResponseCache {
|
|||
pub fn directory(&self) -> Option<&Path> {
|
||||
match self {
|
||||
Self::Disk(cache) => Some(cache.backend().directory()),
|
||||
Self::Memory(_) | Self::Redis { .. } | Self::AzureBlob(_) | Self::RedisSemantic(_) => {
|
||||
None
|
||||
}
|
||||
Self::Memory(_)
|
||||
| Self::Redis { .. }
|
||||
| Self::AzureBlob(_)
|
||||
| Self::Gcs(_)
|
||||
| Self::RedisSemantic(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -209,6 +256,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Memory(cache) => cache.lookup(&request.exact(), now),
|
||||
Self::Redis { cache, .. } => cache.lookup(&request.exact(), now),
|
||||
Self::Gcs(cache) => cache.lookup(&request.exact(), now),
|
||||
Self::Disk(cache) => cache.lookup(&request.exact(), now),
|
||||
Self::AzureBlob(cache) => cache.lookup(&request.exact(), now),
|
||||
Self::RedisSemantic(cache) => cache.lookup(&request.semantic(), now),
|
||||
|
|
@ -224,6 +272,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Memory(cache) => cache.store(&request.exact(), response, now),
|
||||
Self::Redis { cache, .. } => cache.store(&request.exact(), response, now),
|
||||
Self::Gcs(cache) => cache.store(&request.exact(), response, now),
|
||||
Self::Disk(cache) => cache.store(&request.exact(), response, now),
|
||||
Self::AzureBlob(cache) => cache.store(&request.exact(), response, now),
|
||||
Self::RedisSemantic(cache) => cache.store(&request.semantic(), response, now),
|
||||
|
|
@ -238,6 +287,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Memory(cache) => cache.lookup_batch(&Self::exact_requests(requests), now),
|
||||
Self::Redis { cache, .. } => cache.lookup_batch(&Self::exact_requests(requests), now),
|
||||
Self::Gcs(cache) => cache.lookup_batch(&Self::exact_requests(requests), now),
|
||||
Self::Disk(cache) => cache.lookup_batch(&Self::exact_requests(requests), now),
|
||||
Self::AzureBlob(cache) => cache.lookup_batch(&Self::exact_requests(requests), now),
|
||||
Self::RedisSemantic(_) => Err(Error::UnsupportedOperation),
|
||||
|
|
@ -252,6 +302,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Memory(cache) => cache.async_lookup(&request.exact(), now).await,
|
||||
Self::Redis { cache, .. } => cache.async_lookup(&request.exact(), now).await,
|
||||
Self::Gcs(cache) => cache.async_lookup(&request.exact(), now).await,
|
||||
Self::Disk(cache) => cache.async_lookup(&request.exact(), now).await,
|
||||
Self::AzureBlob(cache) => cache.async_lookup(&request.exact(), now).await,
|
||||
Self::RedisSemantic(cache) => cache.async_lookup(&request.semantic(), now).await,
|
||||
|
|
@ -278,6 +329,7 @@ impl NativeResponseCache {
|
|||
.async_store(cache, &request.exact(), response, now)
|
||||
.await
|
||||
}
|
||||
Self::Gcs(cache) => cache.async_store(&request.exact(), response, now).await,
|
||||
Self::Disk(cache) => cache.async_store(&request.exact(), response, now).await,
|
||||
Self::AzureBlob(cache) => cache.async_store(&request.exact(), response, now).await,
|
||||
Self::RedisSemantic(cache) => {
|
||||
|
|
@ -302,6 +354,11 @@ impl NativeResponseCache {
|
|||
.async_lookup_batch(&Self::exact_requests(requests), now)
|
||||
.await
|
||||
}
|
||||
Self::Gcs(cache) => {
|
||||
cache
|
||||
.async_lookup_batch(&Self::exact_requests(requests), now)
|
||||
.await
|
||||
}
|
||||
Self::Disk(cache) => {
|
||||
cache
|
||||
.async_lookup_batch(&Self::exact_requests(requests), now)
|
||||
|
|
@ -344,6 +401,17 @@ impl NativeResponseCache {
|
|||
)
|
||||
.await
|
||||
}
|
||||
Self::Gcs(cache) => {
|
||||
cache
|
||||
.async_store_batch(
|
||||
entries
|
||||
.into_iter()
|
||||
.map(|(request, value)| (request.exact(), value))
|
||||
.collect(),
|
||||
now,
|
||||
)
|
||||
.await
|
||||
}
|
||||
Self::Disk(cache) => {
|
||||
cache
|
||||
.async_store_batch(
|
||||
|
|
@ -389,6 +457,7 @@ impl NativeResponseCache {
|
|||
}
|
||||
cache.async_flush().await
|
||||
}
|
||||
Self::Gcs(cache) => cache.async_flush().await,
|
||||
Self::Disk(cache) => cache.async_flush().await,
|
||||
Self::RedisSemantic(_) => Err(Error::UnsupportedOperation),
|
||||
Self::AzureBlob(cache) => cache.async_flush().await,
|
||||
|
|
@ -399,9 +468,17 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Memory(cache) => cache.test_connection().await,
|
||||
Self::Redis { cache, .. } => cache.test_connection().await,
|
||||
Self::Gcs(cache) => cache.test_connection().await,
|
||||
Self::Disk(cache) => cache.test_connection().await,
|
||||
Self::RedisSemantic(_) => Err(Error::UnsupportedOperation),
|
||||
Self::AzureBlob(cache) => cache.test_connection().await,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn gcs_backend(&self) -> Option<&GcsCache<ResponseCacheCodec>> {
|
||||
match self {
|
||||
Self::Gcs(cache) => Some(cache.backend()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
25
litellm-rust/crates/secrets-hashicorp/Cargo.toml
Normal file
25
litellm-rust/crates/secrets-hashicorp/Cargo.toml
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
[package]
|
||||
name = "litellm-secrets-hashicorp"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-core-utils.workspace = true
|
||||
litellm-secrets-types.workspace = true
|
||||
moka.workspace = true
|
||||
rustify.workspace = true
|
||||
rustify_derive.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio.workspace = true
|
||||
vaultrs.workspace = true
|
||||
veil.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
tempfile = "3"
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
22
litellm-rust/crates/secrets-hashicorp/src/cert_login.rs
Normal file
22
litellm-rust/crates/secrets-hashicorp/src/cert_login.rs
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
#[derive(Debug, rustify_derive::Endpoint)]
|
||||
#[endpoint(path = "/auth/{self.mount}/login", method = "POST")]
|
||||
pub struct CertLoginRequest {
|
||||
#[endpoint(skip)]
|
||||
pub mount: String,
|
||||
#[endpoint(raw)]
|
||||
body: Vec<u8>,
|
||||
}
|
||||
|
||||
impl CertLoginRequest {
|
||||
pub fn new(name: Option<&str>) -> Self {
|
||||
let body: Vec<u8> = match name {
|
||||
Some(name) => serde_json::to_vec(&serde_json::json!({ "name": name }))
|
||||
.expect("json object serialization is infallible"),
|
||||
None => b"{}".to_vec(),
|
||||
};
|
||||
Self {
|
||||
mount: "cert".to_owned(),
|
||||
body,
|
||||
}
|
||||
}
|
||||
}
|
||||
161
litellm-rust/crates/secrets-hashicorp/src/config.rs
Normal file
161
litellm-rust/crates/secrets-hashicorp/src/config.rs
Normal file
|
|
@ -0,0 +1,161 @@
|
|||
use std::{path::PathBuf, time::Duration};
|
||||
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets_types::SecretValue;
|
||||
|
||||
use crate::Error;
|
||||
|
||||
const DEFAULT_ADDRESS: &str = "http://127.0.0.1:8200";
|
||||
const DEFAULT_MOUNT: &str = "secret";
|
||||
const DEFAULT_APPROLE_MOUNT_PATH: &str = "approle";
|
||||
const DEFAULT_REFRESH_INTERVAL: Duration = Duration::from_secs(86400);
|
||||
const HCP_VAULT_ADDR: &str = "HCP_VAULT_ADDR";
|
||||
const HCP_VAULT_TOKEN: &str = "HCP_VAULT_TOKEN";
|
||||
const HCP_VAULT_NAMESPACE: &str = "HCP_VAULT_NAMESPACE";
|
||||
const HCP_VAULT_LOGIN_NAMESPACE: &str = "HCP_VAULT_LOGIN_NAMESPACE";
|
||||
const HCP_VAULT_SECRET_NAMESPACE: &str = "HCP_VAULT_SECRET_NAMESPACE";
|
||||
const HCP_VAULT_MOUNT_NAME: &str = "HCP_VAULT_MOUNT_NAME";
|
||||
const HCP_VAULT_PATH_PREFIX: &str = "HCP_VAULT_PATH_PREFIX";
|
||||
const HCP_VAULT_APPROLE_ROLE_ID: &str = "HCP_VAULT_APPROLE_ROLE_ID";
|
||||
const HCP_VAULT_APPROLE_SECRET_ID: &str = "HCP_VAULT_APPROLE_SECRET_ID";
|
||||
const HCP_VAULT_APPROLE_MOUNT_PATH: &str = "HCP_VAULT_APPROLE_MOUNT_PATH";
|
||||
const HCP_VAULT_CLIENT_CERT: &str = "HCP_VAULT_CLIENT_CERT";
|
||||
const HCP_VAULT_CLIENT_KEY: &str = "HCP_VAULT_CLIENT_KEY";
|
||||
const HCP_VAULT_CERT_ROLE: &str = "HCP_VAULT_CERT_ROLE";
|
||||
const HCP_VAULT_REFRESH_INTERVAL: &str = "HCP_VAULT_REFRESH_INTERVAL";
|
||||
const SECRET_MANAGER_REFRESH_INTERVAL: &str = "SECRET_MANAGER_REFRESH_INTERVAL";
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct AppRoleAuth {
|
||||
pub role_id: String,
|
||||
pub secret_id: SecretValue,
|
||||
pub mount_path: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct TlsCertAuth {
|
||||
pub cert_path: PathBuf,
|
||||
pub key_path: PathBuf,
|
||||
pub role: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct HashicorpVaultConfig {
|
||||
pub address: String,
|
||||
pub token: Option<SecretValue>,
|
||||
pub namespace: Option<String>,
|
||||
pub login_namespace: Option<String>,
|
||||
pub secret_namespace: Option<String>,
|
||||
pub mount: String,
|
||||
pub path_prefix: Option<String>,
|
||||
pub approle: Option<AppRoleAuth>,
|
||||
pub tls_cert: Option<TlsCertAuth>,
|
||||
pub refresh_interval: Duration,
|
||||
}
|
||||
|
||||
impl HashicorpVaultConfig {
|
||||
pub fn from_environment(environment: &dyn Lookup) -> Result<Self, Error> {
|
||||
let address: String = environment
|
||||
.get(HCP_VAULT_ADDR)
|
||||
.and_then(|value| nonempty(value.trim()))
|
||||
.map(|value| value.trim_end_matches('/').to_owned())
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or_else(|| DEFAULT_ADDRESS.to_owned());
|
||||
let token: Option<SecretValue> = environment
|
||||
.get(HCP_VAULT_TOKEN)
|
||||
.and_then(nonempty)
|
||||
.map(SecretValue::new);
|
||||
let namespace: Option<String> = path_component(environment.get(HCP_VAULT_NAMESPACE));
|
||||
let login_namespace: Option<String> =
|
||||
path_component(environment.get(HCP_VAULT_LOGIN_NAMESPACE));
|
||||
let secret_namespace: Option<String> =
|
||||
path_component(environment.get(HCP_VAULT_SECRET_NAMESPACE));
|
||||
let mount: String = path_component(environment.get(HCP_VAULT_MOUNT_NAME))
|
||||
.unwrap_or_else(|| DEFAULT_MOUNT.to_owned());
|
||||
let path_prefix: Option<String> = path_component(environment.get(HCP_VAULT_PATH_PREFIX));
|
||||
let approle: Option<AppRoleAuth> = match (
|
||||
environment
|
||||
.get(HCP_VAULT_APPROLE_ROLE_ID)
|
||||
.and_then(nonempty),
|
||||
environment
|
||||
.get(HCP_VAULT_APPROLE_SECRET_ID)
|
||||
.and_then(nonempty)
|
||||
.map(SecretValue::new),
|
||||
) {
|
||||
(Some(role_id), Some(secret_id)) => Some(AppRoleAuth {
|
||||
role_id,
|
||||
secret_id,
|
||||
mount_path: path_component(environment.get(HCP_VAULT_APPROLE_MOUNT_PATH))
|
||||
.unwrap_or_else(|| DEFAULT_APPROLE_MOUNT_PATH.to_owned()),
|
||||
}),
|
||||
_ => None,
|
||||
};
|
||||
let tls_cert: Option<TlsCertAuth> = match (
|
||||
environment.get(HCP_VAULT_CLIENT_CERT).and_then(nonempty),
|
||||
environment.get(HCP_VAULT_CLIENT_KEY).and_then(nonempty),
|
||||
) {
|
||||
(Some(cert_path), Some(key_path)) => Some(TlsCertAuth {
|
||||
cert_path: PathBuf::from(cert_path),
|
||||
key_path: PathBuf::from(key_path),
|
||||
role: environment.get(HCP_VAULT_CERT_ROLE).and_then(nonempty),
|
||||
}),
|
||||
_ => None,
|
||||
};
|
||||
let refresh_interval: Duration = refresh_interval(environment)?;
|
||||
Ok(Self {
|
||||
address,
|
||||
token,
|
||||
namespace,
|
||||
login_namespace,
|
||||
secret_namespace,
|
||||
mount,
|
||||
path_prefix,
|
||||
approle,
|
||||
tls_cert,
|
||||
refresh_interval,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn login_namespace(&self) -> Option<&str> {
|
||||
self.login_namespace
|
||||
.as_deref()
|
||||
.or(self.namespace.as_deref())
|
||||
}
|
||||
|
||||
pub fn secret_namespace(&self) -> Option<&str> {
|
||||
self.secret_namespace
|
||||
.as_deref()
|
||||
.or(self.namespace.as_deref())
|
||||
}
|
||||
}
|
||||
|
||||
fn nonempty(value: impl AsRef<str>) -> Option<String> {
|
||||
let value: &str = value.as_ref();
|
||||
(!value.is_empty()).then(|| value.to_owned())
|
||||
}
|
||||
|
||||
fn path_component(value: Option<String>) -> Option<String> {
|
||||
value
|
||||
.and_then(|value| nonempty(value.trim()))
|
||||
.map(|value| value.trim_matches('/').to_owned())
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn refresh_interval(environment: &dyn Lookup) -> Result<Duration, Error> {
|
||||
let value: Option<String> = environment
|
||||
.get(HCP_VAULT_REFRESH_INTERVAL)
|
||||
.and_then(nonempty)
|
||||
.or_else(|| {
|
||||
environment
|
||||
.get(SECRET_MANAGER_REFRESH_INTERVAL)
|
||||
.and_then(nonempty)
|
||||
});
|
||||
let Some(value) = value else {
|
||||
return Ok(DEFAULT_REFRESH_INTERVAL);
|
||||
};
|
||||
let seconds: i64 = value.parse().map_err(|_| Error::RefreshInterval)?;
|
||||
if seconds < 0 {
|
||||
return Err(Error::RefreshInterval);
|
||||
}
|
||||
Ok(Duration::from_secs(seconds as u64))
|
||||
}
|
||||
34
litellm-rust/crates/secrets-hashicorp/src/error.rs
Normal file
34
litellm-rust/crates/secrets-hashicorp/src/error.rs
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
#[derive(thiserror::Error, veil::Redact)]
|
||||
pub enum Error {
|
||||
#[error("HashiCorp Vault requires an enterprise license")]
|
||||
EnterpriseRequired,
|
||||
#[error("invalid secret name")]
|
||||
InvalidSecretName(#[from] litellm_secrets_types::Error),
|
||||
#[error("HashiCorp Vault client failed")]
|
||||
Client(
|
||||
#[from]
|
||||
#[redact]
|
||||
vaultrs::error::ClientError,
|
||||
),
|
||||
#[error("HashiCorp Vault client settings are invalid: {message}")]
|
||||
ClientSettings { message: String },
|
||||
#[error("HashiCorp Vault TLS identity could not be configured for {path}: {message}")]
|
||||
TlsIdentity {
|
||||
path: std::path::PathBuf,
|
||||
message: String,
|
||||
},
|
||||
#[error("HashiCorp Vault login returned HTTP {status}")]
|
||||
LoginStatus { status: u16 },
|
||||
#[error("HashiCorp Vault login response is malformed")]
|
||||
MalformedLogin,
|
||||
#[error("HashiCorp Vault authentication is not configured")]
|
||||
NoAuthConfigured,
|
||||
#[error("HashiCorp Vault returned HTTP {status}")]
|
||||
Status { status: u16 },
|
||||
#[error("HashiCorp Vault response payload is malformed")]
|
||||
MalformedPayload,
|
||||
#[error("HashiCorp Vault secret value is not a string")]
|
||||
NonStringValue,
|
||||
#[error("invalid HashiCorp Vault refresh interval")]
|
||||
RefreshInterval,
|
||||
}
|
||||
10
litellm-rust/crates/secrets-hashicorp/src/lib.rs
Normal file
10
litellm-rust/crates/secrets-hashicorp/src/lib.rs
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
#![forbid(unsafe_code)]
|
||||
|
||||
mod cert_login;
|
||||
mod config;
|
||||
mod error;
|
||||
pub mod secret_manager;
|
||||
|
||||
pub use config::{AppRoleAuth, HashicorpVaultConfig, TlsCertAuth};
|
||||
pub use error::Error;
|
||||
pub use secret_manager::{HashicorpVault, SecretLocation};
|
||||
359
litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs
Normal file
359
litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs
Normal file
|
|
@ -0,0 +1,359 @@
|
|||
use std::{
|
||||
collections::HashMap,
|
||||
fmt,
|
||||
sync::Arc,
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets_types::{
|
||||
BaseSecretManager, SecretValue, async_rotate_secret, validate_secret_name,
|
||||
};
|
||||
use moka::future::Cache;
|
||||
use rustify::errors::ClientError as RustifyClientError;
|
||||
use serde_json::Value;
|
||||
use tokio::sync::Mutex;
|
||||
use vaultrs::{
|
||||
api,
|
||||
auth::approle,
|
||||
client::{Identity, VaultClient, VaultClientSettingsBuilder},
|
||||
error::ClientError,
|
||||
kv2,
|
||||
};
|
||||
|
||||
use crate::{Error, HashicorpVaultConfig, TlsCertAuth, cert_login::CertLoginRequest};
|
||||
|
||||
const CACHE_CAPACITY: u64 = 200;
|
||||
|
||||
#[derive(Clone)]
|
||||
struct CachedClient {
|
||||
client: Arc<VaultClient>,
|
||||
expires_at: Option<Instant>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct SecretLocation {
|
||||
pub namespace: Option<String>,
|
||||
pub mount: String,
|
||||
pub path: String,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct HashicorpVault {
|
||||
config: HashicorpVaultConfig,
|
||||
cache: Cache<String, SecretValue>,
|
||||
auth_client: Arc<Mutex<Option<CachedClient>>>,
|
||||
}
|
||||
|
||||
impl fmt::Debug for HashicorpVault {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("HashicorpVault")
|
||||
.field("config", &self.config)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl HashicorpVault {
|
||||
pub fn new(
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
enterprise_enabled: bool,
|
||||
) -> Result<Self, Error> {
|
||||
let config: HashicorpVaultConfig =
|
||||
HashicorpVaultConfig::from_environment(environment.as_ref())?;
|
||||
Self::from_config(config, enterprise_enabled)
|
||||
}
|
||||
|
||||
pub fn from_config(
|
||||
config: HashicorpVaultConfig,
|
||||
enterprise_enabled: bool,
|
||||
) -> Result<Self, Error> {
|
||||
if !enterprise_enabled {
|
||||
return Err(Error::EnterpriseRequired);
|
||||
}
|
||||
let cache: Cache<String, SecretValue> = Cache::builder()
|
||||
.max_capacity(CACHE_CAPACITY)
|
||||
.time_to_live(config.refresh_interval)
|
||||
.build();
|
||||
Ok(Self {
|
||||
config,
|
||||
cache,
|
||||
auth_client: Arc::new(Mutex::new(None)),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn secret_location(&self, secret_name: &str) -> Result<SecretLocation, Error> {
|
||||
validate_secret_name(secret_name).map_err(Error::InvalidSecretName)?;
|
||||
let path: String = [
|
||||
self.config.path_prefix.clone(),
|
||||
Some(secret_name.to_owned()),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.collect::<Vec<String>>()
|
||||
.join("/");
|
||||
Ok(SecretLocation {
|
||||
namespace: self.config.secret_namespace().map(str::to_owned),
|
||||
mount: self.config.mount.clone(),
|
||||
path,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn config(&self) -> &HashicorpVaultConfig {
|
||||
&self.config
|
||||
}
|
||||
|
||||
pub async fn async_read_secret(&self, secret_name: &str) -> Result<Option<SecretValue>, Error> {
|
||||
let location: SecretLocation = self.secret_location(secret_name)?;
|
||||
let cache_key: String = cache_key(&location);
|
||||
if let Some(value) = self.cache.get(&cache_key).await {
|
||||
return Ok(Some(value));
|
||||
}
|
||||
let client: Arc<VaultClient> = self.vault_client().await?;
|
||||
let data: HashMap<String, Value> =
|
||||
match kv2::read(client.as_ref(), &location.mount, &location.path).await {
|
||||
Ok(data) => data,
|
||||
Err(error) if api_status(&error) == Some(404) => return Ok(None),
|
||||
Err(error) => return Err(map_api_error(error, ErrorContext::Read)),
|
||||
};
|
||||
let Some(value) = data.get("key") else {
|
||||
return Ok(None);
|
||||
};
|
||||
let value: &str = value.as_str().ok_or(Error::NonStringValue)?;
|
||||
let value: SecretValue = SecretValue::new(value);
|
||||
self.cache.insert(cache_key, value.clone()).await;
|
||||
Ok(Some(value))
|
||||
}
|
||||
|
||||
pub async fn async_write_secret(
|
||||
&self,
|
||||
secret_name: &str,
|
||||
value: SecretValue,
|
||||
description: Option<&str>,
|
||||
) -> Result<Value, Error> {
|
||||
let location: SecretLocation = self.secret_location(secret_name)?;
|
||||
let cache_key: String = cache_key(&location);
|
||||
let data: HashMap<String, Value> = match description {
|
||||
Some(description) => [
|
||||
("key".to_owned(), Value::String(value.expose().to_owned())),
|
||||
(
|
||||
"description".to_owned(),
|
||||
Value::String(description.to_owned()),
|
||||
),
|
||||
]
|
||||
.into_iter()
|
||||
.collect(),
|
||||
None => [("key".to_owned(), Value::String(value.expose().to_owned()))]
|
||||
.into_iter()
|
||||
.collect(),
|
||||
};
|
||||
let client: Arc<VaultClient> = self.vault_client().await?;
|
||||
let metadata = kv2::set(client.as_ref(), &location.mount, &location.path, &data)
|
||||
.await
|
||||
.map_err(|error| map_api_error(error, ErrorContext::Secret))?;
|
||||
self.cache.invalidate(&cache_key).await;
|
||||
serde_json::to_value(metadata)
|
||||
.map_err(|source| Error::Client(ClientError::JsonParseError { source }))
|
||||
}
|
||||
|
||||
pub async fn async_delete_secret(&self, secret_name: &str) -> Result<(), Error> {
|
||||
let location: SecretLocation = self.secret_location(secret_name)?;
|
||||
let cache_key: String = cache_key(&location);
|
||||
let client: Arc<VaultClient> = self.vault_client().await?;
|
||||
kv2::delete_latest(client.as_ref(), &location.mount, &location.path)
|
||||
.await
|
||||
.map_err(|error| map_api_error(error, ErrorContext::Secret))?;
|
||||
self.cache.invalidate(&cache_key).await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn async_rotate_secret(
|
||||
&self,
|
||||
current_name: &str,
|
||||
new_name: &str,
|
||||
value: &SecretValue,
|
||||
) -> Result<Value, Error> {
|
||||
async_rotate_secret(self, current_name, new_name, value).await
|
||||
}
|
||||
|
||||
async fn vault_client(&self) -> Result<Arc<VaultClient>, Error> {
|
||||
let mut cached = self.auth_client.lock().await;
|
||||
if let Some(entry) = cached.as_ref()
|
||||
&& entry
|
||||
.expires_at
|
||||
.is_none_or(|expires_at| expires_at > Instant::now())
|
||||
{
|
||||
return Ok(entry.client.clone());
|
||||
}
|
||||
|
||||
let (client, expires_at): (VaultClient, Option<Instant>) =
|
||||
match (self.config.approle.as_ref(), self.config.tls_cert.as_ref()) {
|
||||
(Some(approle), _) => {
|
||||
let login_client: VaultClient =
|
||||
self.build_client(self.config.login_namespace(), "")?;
|
||||
let auth = approle::login(
|
||||
&login_client,
|
||||
&approle.mount_path,
|
||||
&approle.role_id,
|
||||
approle.secret_id.expose(),
|
||||
)
|
||||
.await
|
||||
.map_err(|error| map_api_error(error, ErrorContext::Login))?;
|
||||
(
|
||||
self.build_client(self.config.secret_namespace(), &auth.client_token)?,
|
||||
token_expiry(auth.lease_duration),
|
||||
)
|
||||
}
|
||||
(None, Some(tls)) => {
|
||||
let login_client: VaultClient =
|
||||
self.build_client(self.config.login_namespace(), "")?;
|
||||
let endpoint: CertLoginRequest = CertLoginRequest::new(tls.role.as_deref());
|
||||
let auth = api::auth(&login_client, endpoint)
|
||||
.await
|
||||
.map_err(|error| map_api_error(error, ErrorContext::Login))?;
|
||||
(
|
||||
self.build_client(self.config.secret_namespace(), &auth.client_token)?,
|
||||
token_expiry(auth.lease_duration),
|
||||
)
|
||||
}
|
||||
(None, None) => {
|
||||
let token: SecretValue =
|
||||
self.config.token.clone().ok_or(Error::NoAuthConfigured)?;
|
||||
(
|
||||
self.build_client(self.config.secret_namespace(), token.expose())?,
|
||||
None,
|
||||
)
|
||||
}
|
||||
};
|
||||
let client: Arc<VaultClient> = Arc::new(client);
|
||||
*cached = Some(CachedClient {
|
||||
client: client.clone(),
|
||||
expires_at,
|
||||
});
|
||||
Ok(client)
|
||||
}
|
||||
|
||||
fn build_client(&self, namespace: Option<&str>, token: &str) -> Result<VaultClient, Error> {
|
||||
let settings = VaultClientSettingsBuilder::default()
|
||||
.address(&self.config.address)
|
||||
.token(token.to_owned())
|
||||
.namespace(namespace.map(str::to_owned))
|
||||
.identity(identity_for(self.config.tls_cert.as_ref())?)
|
||||
.ca_certs(Vec::new())
|
||||
.verify(true)
|
||||
.build()
|
||||
.map_err(|message| Error::ClientSettings {
|
||||
message: message.to_string(),
|
||||
})?;
|
||||
VaultClient::new(settings).map_err(Error::Client)
|
||||
}
|
||||
}
|
||||
|
||||
impl BaseSecretManager for HashicorpVault {
|
||||
type Error = Error;
|
||||
type WriteResponse = Value;
|
||||
type DeleteResponse = ();
|
||||
|
||||
async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
|
||||
HashicorpVault::async_read_secret(self, name).await
|
||||
}
|
||||
|
||||
async fn async_write_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
value: &SecretValue,
|
||||
description: Option<&str>,
|
||||
) -> Result<Value, Error> {
|
||||
HashicorpVault::async_write_secret(self, name, value.clone(), description).await
|
||||
}
|
||||
|
||||
async fn async_delete_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
_recovery_window_in_days: i64,
|
||||
) -> Result<(), Error> {
|
||||
HashicorpVault::async_delete_secret(self, name).await
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum ErrorContext {
|
||||
Login,
|
||||
Read,
|
||||
Secret,
|
||||
}
|
||||
|
||||
fn cache_key(location: &SecretLocation) -> String {
|
||||
format!(
|
||||
"{:?}/{}/{}",
|
||||
location.namespace, location.mount, location.path
|
||||
)
|
||||
}
|
||||
|
||||
fn identity_for(tls: Option<&TlsCertAuth>) -> Result<Option<Identity>, Error> {
|
||||
tls.map(|tls| {
|
||||
let cert: Vec<u8> = std::fs::read(&tls.cert_path).map_err(|source| Error::TlsIdentity {
|
||||
path: tls.cert_path.clone(),
|
||||
message: source.to_string(),
|
||||
})?;
|
||||
let key: Vec<u8> = std::fs::read(&tls.key_path).map_err(|source| Error::TlsIdentity {
|
||||
path: tls.key_path.clone(),
|
||||
message: source.to_string(),
|
||||
})?;
|
||||
Identity::from_pem(&[cert.as_slice(), key.as_slice()].concat()).map_err(|source| {
|
||||
Error::TlsIdentity {
|
||||
path: tls.cert_path.clone(),
|
||||
message: source.to_string(),
|
||||
}
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn map_api_error(error: ClientError, context: ErrorContext) -> Error {
|
||||
match error {
|
||||
ClientError::APIError { code, .. } => match context {
|
||||
ErrorContext::Login => Error::LoginStatus { status: code },
|
||||
ErrorContext::Read | ErrorContext::Secret => Error::Status { status: code },
|
||||
},
|
||||
ClientError::JsonParseError { source } => match context {
|
||||
ErrorContext::Login => Error::MalformedLogin,
|
||||
ErrorContext::Read => Error::MalformedPayload,
|
||||
ErrorContext::Secret => Error::Client(ClientError::JsonParseError { source }),
|
||||
},
|
||||
ClientError::ResponseEmptyError | ClientError::ResponseDataEmptyError => {
|
||||
malformed_response(context)
|
||||
}
|
||||
ClientError::RestClientError { source } => match source {
|
||||
RustifyClientError::ServerResponseError { code, .. } => match context {
|
||||
ErrorContext::Login => Error::LoginStatus { status: code },
|
||||
ErrorContext::Read | ErrorContext::Secret => Error::Status { status: code },
|
||||
},
|
||||
RustifyClientError::ResponseParseError { .. } => malformed_response(context),
|
||||
source => Error::Client(ClientError::RestClientError { source }),
|
||||
},
|
||||
error => Error::Client(error),
|
||||
}
|
||||
}
|
||||
|
||||
fn api_status(error: &ClientError) -> Option<u16> {
|
||||
match error {
|
||||
ClientError::APIError { code, .. } => Some(*code),
|
||||
ClientError::RestClientError {
|
||||
source: RustifyClientError::ServerResponseError { code, .. },
|
||||
} => Some(*code),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn malformed_response(context: ErrorContext) -> Error {
|
||||
match context {
|
||||
ErrorContext::Login => Error::MalformedLogin,
|
||||
ErrorContext::Read => Error::MalformedPayload,
|
||||
ErrorContext::Secret => Error::Client(ClientError::ResponseDataEmptyError),
|
||||
}
|
||||
}
|
||||
|
||||
fn token_expiry(lease_duration: u64) -> Option<Instant> {
|
||||
(lease_duration > 0).then(|| Instant::now() + Duration::from_secs(lease_duration))
|
||||
}
|
||||
602
litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs
Normal file
602
litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs
Normal file
|
|
@ -0,0 +1,602 @@
|
|||
use std::{collections::HashMap, sync::Arc, time::Duration};
|
||||
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets_hashicorp::{Error, HashicorpVault, HashicorpVaultConfig};
|
||||
use litellm_secrets_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
use serde_json::json;
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{body_json, header, method, path},
|
||||
};
|
||||
|
||||
fn config(server: &MockServer, values: &[(&str, &str)]) -> HashicorpVaultConfig {
|
||||
let mut environment_values: HashMap<String, String> = values
|
||||
.iter()
|
||||
.map(|(name, value)| ((*name).to_owned(), (*value).to_owned()))
|
||||
.collect();
|
||||
environment_values.insert("HCP_VAULT_ADDR".to_owned(), server.uri());
|
||||
let environment: Arc<dyn Lookup + Send + Sync> =
|
||||
Arc::new(move |name: &str| environment_values.get(name).cloned());
|
||||
HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap()
|
||||
}
|
||||
|
||||
fn manager(server: &MockServer, values: &[(&str, &str)]) -> HashicorpVault {
|
||||
HashicorpVault::from_config(config(server, values), true).unwrap()
|
||||
}
|
||||
|
||||
fn auth_response(token: &str, lease_duration: u64) -> serde_json::Value {
|
||||
json!({
|
||||
"auth": {
|
||||
"client_token": token,
|
||||
"accessor": "",
|
||||
"policies": [],
|
||||
"token_policies": [],
|
||||
"metadata": null,
|
||||
"lease_duration": lease_duration,
|
||||
"renewable": false,
|
||||
"entity_id": "",
|
||||
"token_type": "service",
|
||||
"orphan": false
|
||||
},
|
||||
"lease_id": "",
|
||||
"lease_duration": lease_duration,
|
||||
"renewable": false,
|
||||
"request_id": "",
|
||||
"warnings": null,
|
||||
"wrap_info": null
|
||||
})
|
||||
}
|
||||
|
||||
fn read_response(data: serde_json::Value) -> serde_json::Value {
|
||||
json!({
|
||||
"data": {
|
||||
"data": data,
|
||||
"metadata": {
|
||||
"created_time": "",
|
||||
"deletion_time": "",
|
||||
"custom_metadata": null,
|
||||
"destroyed": false,
|
||||
"version": 1
|
||||
}
|
||||
},
|
||||
"lease_id": "",
|
||||
"lease_duration": 0,
|
||||
"renewable": false,
|
||||
"request_id": "",
|
||||
"warnings": null,
|
||||
"wrap_info": null
|
||||
})
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn token_reads_use_vault_headers_and_cache_values() {
|
||||
let server: MockServer = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/v1/secret/data/name"))
|
||||
.and(header("X-Vault-Token", "token"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager: HashicorpVault = manager(&server, &[("HCP_VAULT_TOKEN", "token")]);
|
||||
|
||||
assert_eq!(
|
||||
manager
|
||||
.async_read_secret("name")
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"value"
|
||||
);
|
||||
let requests = server.received_requests().await.unwrap();
|
||||
assert!(
|
||||
requests
|
||||
.iter()
|
||||
.all(|request| !request.headers.contains_key("X-Vault-Namespace"))
|
||||
);
|
||||
assert_eq!(
|
||||
manager
|
||||
.async_read_secret("name")
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"value"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn namespace_mount_and_prefix_are_sanitized_in_the_url() {
|
||||
let server: MockServer = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/v1/kv-prod/data/virtual-keys/name"))
|
||||
.and(header("X-Vault-Namespace", "team-a"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager: HashicorpVault = manager(
|
||||
&server,
|
||||
&[
|
||||
("HCP_VAULT_TOKEN", "token"),
|
||||
("HCP_VAULT_SECRET_NAMESPACE", " /team-a/ "),
|
||||
("HCP_VAULT_MOUNT_NAME", " /kv-prod/ "),
|
||||
("HCP_VAULT_PATH_PREFIX", " /virtual-keys/ "),
|
||||
],
|
||||
);
|
||||
|
||||
let location = manager.secret_location("name").unwrap();
|
||||
assert_eq!(location.namespace.as_deref(), Some("team-a"));
|
||||
assert_eq!(location.mount, "kv-prod");
|
||||
assert_eq!(location.path, "virtual-keys/name");
|
||||
assert!(manager.async_read_secret("name").await.unwrap().is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trailing_address_slashes_are_removed() {
|
||||
let environment: Arc<dyn Lookup + Send + Sync> = Arc::new(|name: &str| match name {
|
||||
"HCP_VAULT_ADDR" => Some("http://vault.test:8200///".to_owned()),
|
||||
"HCP_VAULT_TOKEN" => Some("token".to_owned()),
|
||||
_ => None,
|
||||
});
|
||||
let config: HashicorpVaultConfig =
|
||||
HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap();
|
||||
let manager: HashicorpVault = HashicorpVault::from_config(config, true).unwrap();
|
||||
|
||||
assert_eq!(
|
||||
manager.secret_location("name").unwrap(),
|
||||
litellm_secrets_hashicorp::SecretLocation {
|
||||
namespace: None,
|
||||
mount: "secret".to_owned(),
|
||||
path: "name".to_owned(),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case("-1")]
|
||||
#[case("not-a-number")]
|
||||
fn invalid_refresh_intervals_are_rejected(#[case] value: &str) {
|
||||
let environment: Arc<dyn Lookup + Send + Sync> = Arc::new(move |name: &str| match name {
|
||||
"HCP_VAULT_REFRESH_INTERVAL" => Some(value.to_owned()),
|
||||
_ => None,
|
||||
});
|
||||
|
||||
assert!(matches!(
|
||||
HashicorpVaultConfig::from_environment(environment.as_ref()),
|
||||
Err(Error::RefreshInterval)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn approle_login_uses_namespace_and_reuses_the_token() {
|
||||
let server: MockServer = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/v1/auth/custom-approle/login"))
|
||||
.and(header("X-Vault-Namespace", "login-root"))
|
||||
.and(body_json(json!({"role_id": "role", "secret_id": "secret"})))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(auth_response("login-token", 3600)))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/v1/secret/data/name"))
|
||||
.and(header("X-Vault-Token", "login-token"))
|
||||
.and(header("X-Vault-Namespace", "secret-root"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/v1/secret/data/name-2"))
|
||||
.respond_with(ResponseTemplate::new(404).set_body_json(json!({"errors": ["missing"]})))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager: HashicorpVault = manager(
|
||||
&server,
|
||||
&[
|
||||
("HCP_VAULT_APPROLE_ROLE_ID", "role"),
|
||||
("HCP_VAULT_APPROLE_SECRET_ID", "secret"),
|
||||
("HCP_VAULT_APPROLE_MOUNT_PATH", "custom-approle"),
|
||||
("HCP_VAULT_NAMESPACE", "secret-root"),
|
||||
("HCP_VAULT_LOGIN_NAMESPACE", "login-root"),
|
||||
],
|
||||
);
|
||||
|
||||
assert!(manager.async_read_secret("name").await.unwrap().is_some());
|
||||
assert!(manager.async_read_secret("name-2").await.unwrap().is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn approle_tokens_expire_after_the_vault_lease() {
|
||||
let server: MockServer = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/v1/auth/approle/login"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(auth_response("login-token", 1)))
|
||||
.expect(2)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))),
|
||||
)
|
||||
.expect(2)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager: HashicorpVault = manager(
|
||||
&server,
|
||||
&[
|
||||
("HCP_VAULT_APPROLE_ROLE_ID", "role"),
|
||||
("HCP_VAULT_APPROLE_SECRET_ID", "secret"),
|
||||
("HCP_VAULT_REFRESH_INTERVAL", "0"),
|
||||
],
|
||||
);
|
||||
|
||||
assert!(manager.async_read_secret("first").await.unwrap().is_some());
|
||||
tokio::time::sleep(Duration::from_secs(1) + Duration::from_millis(50)).await;
|
||||
assert!(manager.async_read_secret("second").await.unwrap().is_some());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tls_login_posts_the_role_and_uses_the_client_identity() {
|
||||
let server: MockServer = MockServer::start().await;
|
||||
let directory: tempfile::TempDir = tempfile::tempdir().unwrap();
|
||||
let cert_path = directory.path().join("client.crt");
|
||||
let key_path = directory.path().join("client.key");
|
||||
std::fs::write(&cert_path, TEST_CERTIFICATE).unwrap();
|
||||
std::fs::write(&key_path, TEST_PRIVATE_KEY).unwrap();
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/v1/auth/cert/login"))
|
||||
.and(header("X-Vault-Namespace", "login-ns"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(auth_response("cert-token", 0)))
|
||||
.expect(2)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/v1/secret/data/name"))
|
||||
.and(header("X-Vault-Token", "cert-token"))
|
||||
.and(header("X-Vault-Namespace", "secret-ns"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))),
|
||||
)
|
||||
.expect(2)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let role_values: HashMap<String, String> = HashMap::from([
|
||||
("HCP_VAULT_ADDR".to_owned(), server.uri()),
|
||||
(
|
||||
"HCP_VAULT_CLIENT_CERT".to_owned(),
|
||||
cert_path.to_str().unwrap().to_owned(),
|
||||
),
|
||||
(
|
||||
"HCP_VAULT_CLIENT_KEY".to_owned(),
|
||||
key_path.to_str().unwrap().to_owned(),
|
||||
),
|
||||
("HCP_VAULT_CERT_ROLE".to_owned(), "vault-role".to_owned()),
|
||||
(
|
||||
"HCP_VAULT_LOGIN_NAMESPACE".to_owned(),
|
||||
"login-ns".to_owned(),
|
||||
),
|
||||
(
|
||||
"HCP_VAULT_SECRET_NAMESPACE".to_owned(),
|
||||
"secret-ns".to_owned(),
|
||||
),
|
||||
]);
|
||||
let role_environment: Arc<dyn Lookup + Send + Sync> =
|
||||
Arc::new(move |name: &str| role_values.get(name).cloned());
|
||||
let role_manager: HashicorpVault = HashicorpVault::new(role_environment, true).unwrap();
|
||||
assert!(
|
||||
role_manager
|
||||
.async_read_secret("name")
|
||||
.await
|
||||
.unwrap()
|
||||
.is_some()
|
||||
);
|
||||
|
||||
let no_role_values: HashMap<String, String> = HashMap::from([
|
||||
("HCP_VAULT_ADDR".to_owned(), server.uri()),
|
||||
(
|
||||
"HCP_VAULT_CLIENT_CERT".to_owned(),
|
||||
cert_path.to_str().unwrap().to_owned(),
|
||||
),
|
||||
(
|
||||
"HCP_VAULT_CLIENT_KEY".to_owned(),
|
||||
key_path.to_str().unwrap().to_owned(),
|
||||
),
|
||||
(
|
||||
"HCP_VAULT_LOGIN_NAMESPACE".to_owned(),
|
||||
"login-ns".to_owned(),
|
||||
),
|
||||
(
|
||||
"HCP_VAULT_SECRET_NAMESPACE".to_owned(),
|
||||
"secret-ns".to_owned(),
|
||||
),
|
||||
]);
|
||||
let no_role_environment: Arc<dyn Lookup + Send + Sync> =
|
||||
Arc::new(move |name: &str| no_role_values.get(name).cloned());
|
||||
let no_role_manager: HashicorpVault = HashicorpVault::new(no_role_environment, true).unwrap();
|
||||
assert!(
|
||||
no_role_manager
|
||||
.async_read_secret("name")
|
||||
.await
|
||||
.unwrap()
|
||||
.is_some()
|
||||
);
|
||||
let login_bodies: Vec<serde_json::Value> = server
|
||||
.received_requests()
|
||||
.await
|
||||
.unwrap()
|
||||
.iter()
|
||||
.filter(|request| request.method.as_str() == "POST")
|
||||
.map(|request| serde_json::from_slice(&request.body).unwrap())
|
||||
.collect();
|
||||
assert!(login_bodies.contains(&json!({"name": "vault-role"})));
|
||||
assert!(login_bodies.contains(&json!({})));
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::missing(404, json!({"errors": ["missing"]}), 0)]
|
||||
#[case::malformed(200, json!({"data": "invalid"}), 1)]
|
||||
#[case::missing_key(200, json!({}), 0)]
|
||||
#[case::non_string(200, json!({"key": 1}), 2)]
|
||||
#[tokio::test]
|
||||
async fn read_responses_distinguish_absence_and_malformed_payloads(
|
||||
#[case] status: u16,
|
||||
#[case] body: serde_json::Value,
|
||||
#[case] expected: u8,
|
||||
) {
|
||||
let server: MockServer = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.respond_with(ResponseTemplate::new(status).set_body_json(
|
||||
if status == 200 && expected != 1 {
|
||||
read_response(body)
|
||||
} else {
|
||||
body
|
||||
},
|
||||
))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let result: Result<Option<SecretValue>, Error> =
|
||||
manager(&server, &[("HCP_VAULT_TOKEN", "token")])
|
||||
.async_read_secret("name")
|
||||
.await;
|
||||
match expected {
|
||||
0 => assert!(result.unwrap().is_none()),
|
||||
1 => assert!(matches!(result, Err(Error::MalformedPayload))),
|
||||
2 => assert!(matches!(result, Err(Error::NonStringValue))),
|
||||
_ => unreachable!(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn write_and_delete_invalidate_the_read_cache() {
|
||||
let server: MockServer = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/v1/secret/data/name"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))),
|
||||
)
|
||||
.expect(2)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/v1/secret/data/name"))
|
||||
.and(body_json(
|
||||
json!({"data": {"key": "updated", "description": "description"}}),
|
||||
))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||
"data": {
|
||||
"created_time": "",
|
||||
"deletion_time": "",
|
||||
"custom_metadata": null,
|
||||
"destroyed": false,
|
||||
"version": 2
|
||||
},
|
||||
"lease_id": "",
|
||||
"lease_duration": 0,
|
||||
"renewable": false,
|
||||
"request_id": "",
|
||||
"warnings": null,
|
||||
"wrap_info": null
|
||||
})))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("DELETE"))
|
||||
.and(path("/v1/secret/data/name"))
|
||||
.respond_with(ResponseTemplate::new(204))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager: HashicorpVault = manager(&server, &[("HCP_VAULT_TOKEN", "token")]);
|
||||
|
||||
assert!(manager.async_read_secret("name").await.unwrap().is_some());
|
||||
assert!(
|
||||
manager
|
||||
.async_write_secret("name", SecretValue::new("updated"), Some("description"))
|
||||
.await
|
||||
.is_ok()
|
||||
);
|
||||
assert!(manager.async_read_secret("name").await.unwrap().is_some());
|
||||
manager.async_delete_secret("name").await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn no_auth_and_invalid_names_fail_without_requests() {
|
||||
let server: MockServer = MockServer::start().await;
|
||||
let manager: HashicorpVault = manager(&server, &[]);
|
||||
|
||||
assert!(matches!(
|
||||
manager.async_read_secret("name").await,
|
||||
Err(Error::NoAuthConfigured)
|
||||
));
|
||||
assert!(matches!(
|
||||
manager.async_read_secret("../name").await,
|
||||
Err(Error::InvalidSecretName(_))
|
||||
));
|
||||
assert!(server.received_requests().await.unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn debug_output_redacts_authentication_values() {
|
||||
let server: MockServer = MockServer::start().await;
|
||||
let manager: HashicorpVault =
|
||||
HashicorpVault::from_config(config(&server, &[("HCP_VAULT_TOKEN", "token-value")]), true)
|
||||
.unwrap();
|
||||
let debug: String = format!("{manager:?}");
|
||||
assert!(!debug.contains("token-value"));
|
||||
assert!(!debug.contains("secret-id"));
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct ParityCase {
|
||||
env: HashMap<String, String>,
|
||||
expected_secret_url: String,
|
||||
expected_login_url: Option<String>,
|
||||
expected_login_namespace: Option<String>,
|
||||
expected_secret_namespace: Option<String>,
|
||||
secret_name: String,
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn configuration_matches_python_parity_fixture() {
|
||||
let cases: Vec<ParityCase> = serde_json::from_str(include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/../../../tests/test_litellm/secret_managers/hashicorp_vault_parity.json"
|
||||
)))
|
||||
.unwrap();
|
||||
for case in cases {
|
||||
let values: HashMap<String, String> = case.env.clone();
|
||||
let environment: Arc<dyn Lookup + Send + Sync> =
|
||||
Arc::new(move |name: &str| values.get(name).cloned());
|
||||
let config: HashicorpVaultConfig =
|
||||
HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap();
|
||||
let manager: HashicorpVault = HashicorpVault::from_config(config.clone(), true).unwrap();
|
||||
let location = manager.secret_location(&case.secret_name).unwrap();
|
||||
let namespace = location
|
||||
.namespace
|
||||
.as_deref()
|
||||
.map(|namespace| format!("{namespace}/"))
|
||||
.unwrap_or_default();
|
||||
assert_eq!(
|
||||
format!(
|
||||
"{}/v1/{}{}/data/{}",
|
||||
config.address, namespace, location.mount, location.path
|
||||
),
|
||||
case.expected_secret_url
|
||||
);
|
||||
let login_url = config.approle.as_ref().map_or_else(
|
||||
|| {
|
||||
config
|
||||
.tls_cert
|
||||
.as_ref()
|
||||
.map(|_| format!("{}/v1/auth/cert/login", config.address))
|
||||
},
|
||||
|approle| {
|
||||
Some(format!(
|
||||
"{}/v1/auth/{}/login",
|
||||
config.address, approle.mount_path
|
||||
))
|
||||
},
|
||||
);
|
||||
assert_eq!(login_url, case.expected_login_url);
|
||||
assert_eq!(
|
||||
manager.config().login_namespace(),
|
||||
case.expected_login_namespace.as_deref()
|
||||
);
|
||||
assert_eq!(
|
||||
manager.config().secret_namespace(),
|
||||
case.expected_secret_namespace.as_deref()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore]
|
||||
async fn live_vault_round_trip() {
|
||||
let environment: Arc<dyn Lookup + Send + Sync> =
|
||||
Arc::new(litellm_core_utils::settings::ProcessEnvironment);
|
||||
let manager: HashicorpVault = HashicorpVault::new(environment, true).unwrap();
|
||||
let name: String = std::env::var("LITELLM_VAULT_LIVE_SECRET_NAME").unwrap();
|
||||
let value: SecretValue = SecretValue::new("native-live-value");
|
||||
let location = manager.secret_location(&name).unwrap();
|
||||
println!(
|
||||
"native provenance: {} vaultrs {} {:?} {} {}",
|
||||
module_path!(),
|
||||
manager.config().address,
|
||||
location.namespace,
|
||||
location.mount,
|
||||
location.path
|
||||
);
|
||||
manager
|
||||
.async_write_secret(&name, value.clone(), None)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
manager.async_read_secret(&name).await.unwrap().unwrap(),
|
||||
value
|
||||
);
|
||||
manager.async_delete_secret(&name).await.unwrap();
|
||||
assert!(manager.async_read_secret(&name).await.unwrap().is_none());
|
||||
}
|
||||
|
||||
const TEST_CERTIFICATE: &str = "-----BEGIN CERTIFICATE-----
|
||||
MIIDDzCCAfegAwIBAgIUeMzLFLM/mRbPGbNAew5N2UTscocwDQYJKoZIhvcNAQEL
|
||||
BQAwFzEVMBMGA1UEAwwMbGl0ZWxsbS10ZXN0MB4XDTI2MDkyMTIwMjA1OVoXDTI2
|
||||
MDkyMjIwMjA1OVowFzEVMBMGA1UEAwwMbGl0ZWxsbS10ZXN0MIIBIjANBgkqhkiG
|
||||
9w0BAQEFAAOCAQ8AMIIBCgKCAQEAveYoSUJXybmkHmQsBfhBcv2Ob5Oy8ejZu+B3
|
||||
vTnrPumW4ANi1XXKBSazRGB3fEtAgr+3KhKeHaSKEQeBwJkAEBfdmQv0tpXICwHs
|
||||
1kFNtU0owy54HVW5/ia+LMszsFcPzVIoMnbUOuiKr9RaV7P+IEFzILPBVuV4DoYH
|
||||
yocjD3+9QNqokWgNL8LK37JijmNEFVaKFz0X6SyL2VRDlfPWTEBK52Gp/pvDgA6G
|
||||
eTSfyI+kCm9h5ECTYUAtmatk9WPVS8sWOqV1EXVanFyYBU+mDxoywAS1/6CHeIPh
|
||||
bNmCOZjPoO9qWBJ7ZyGhOconBigXY8qnlXymev+44IPHrx4urwIDAQABo1MwUTAd
|
||||
BgNVHQ4EFgQUvaZrZ6HKtbr3ekeZmgy4b5Pq95QwHwYDVR0jBBgwFoAUvaZrZ6HK
|
||||
tbr3ekeZmgy4b5Pq95QwDwYDVR0TAQH/BAUwAwEB/zANBgkqhkiG9w0BAQsFAAOC
|
||||
AQEAEejrD8d1qDxW55XxQ4IC31rufoEvDV955jyvh2kALPaN/i5oWsBGI+UAQZna
|
||||
aaoQXwzlmHrtDUBWl0LztVTUamIleUep2+PLLauqqt43vxppxMX8Jn2mnPO20YE/
|
||||
hIzGx0jN/LBG8PDyLSvHdlgjP9ofA4Vg4rTQugdXRgOvlCE/epnH/MADcg9KYJtJ
|
||||
C1RObCIkL3LcdUbjStJRCY/U/FeWcgyncEPz95OFDkbrlNDajb6o6CkYfouqvhTc
|
||||
8XlgjjAVKIbAbRgbVu3elsquuFM97x2DzWDjkrMNmDt1FJ9ubK36gL6B3o0UMaoQ
|
||||
00R7x/eqvH+EkWa/2ekW9lpleQ==
|
||||
-----END CERTIFICATE-----
|
||||
";
|
||||
|
||||
const TEST_PRIVATE_KEY: &str = "-----BEGIN PRIVATE KEY-----
|
||||
MIIEvQIBADANBgkqhkiG9w0BAQEFAASCBKcwggSjAgEAAoIBAQC95ihJQlfJuaQe
|
||||
ZCwF+EFy/Y5vk7Lx6Nm74He9Oes+6ZbgA2LVdcoFJrNEYHd8S0CCv7cqEp4dpIoR
|
||||
B4HAmQAQF92ZC/S2lcgLAezWQU21TSjDLngdVbn+Jr4syzOwVw/NUigydtQ66Iqv
|
||||
1FpXs/4gQXMgs8FW5XgOhgfKhyMPf71A2qiRaA0vwsrfsmKOY0QVVooXPRfpLIvZ
|
||||
VEOV89ZMQErnYan+m8OADoZ5NJ/Ij6QKb2HkQJNhQC2Zq2T1Y9VLyxY6pXURdVqc
|
||||
XJgFT6YPGjLABLX/oId4g+Fs2YI5mM+g72pYEntnIaE5yicGKBdjyqeVfKZ6/7jg
|
||||
g8evHi6vAgMBAAECggEAGdJjlP6b8Fa5bdaCM/ebcrbuuNZVJVbb0JPHxGfNSLs7
|
||||
pE9hj5QaOdQW2Uviw3h6F61ZCzQH4xD+Iy2po5ZKb2XHYKnDB1bboj+LRGER337T
|
||||
9aJqe9at2VTMVEv3Rdm40NsEk0QcPLxlK16NQFK90gYEUSSQPDAswJDSG2R/zHn+
|
||||
vADI907mW/goEJHeLn8PWGlNlSiR6x+5JJtq+GXCzUzVvJYQSCLGxCSl2x2H+0g7
|
||||
NhFI0zPpdzNmO/h+yhzaFb6Rp5U8+ZsnZ3qYjQ/03gw1myTDKJt1YaO9JvArnNYX
|
||||
hcJQQ8Rt0bHhcrZA16bBOpqZlo5pKCicwI/netgN8QKBgQDcFz7AzdJ26sMSV32V
|
||||
rwrMgIoggt8qDjO1ARwqW35A1TIge0FoW4M4KpsXQGGfT341uU1esXEcyZ/1L/5X
|
||||
3ql2gX4DbOYLZLWYzZGR2hq33oi8HkhN98QrEwL9emSH8NqYX3Xxja3PrmCrSYJe
|
||||
Zbnd9TIm2XkxyMoyXJu6M/QvnwKBgQDc4dzqTbxoGEGa5MuJoGmMwPnqgdG9UM5J
|
||||
eExVnh7osxc2sOdsiPeRjjQTxs9v2kJwctC359OJoo9yGaaJeSghU4LEWJo1sqnA
|
||||
fzSCLammYvtVAtniyNv5Mxk/6Uimi4NNDKaAKB+m4K2uSn3U9AmY7KPYMGaSbS9W
|
||||
XSnobjxm8QKBgC8bPpAvvWs8ZhIn7bY659nLbUT2HeO3dHO6UBf0yzn/J6JyHxbB
|
||||
93zvCZDZc8uQTRgcmCW7XtVlhjoJUqvl+Wlm39zF0xr/LCsPXKfWAb/2/lcdOCaP
|
||||
8Emz4QD10EyUTYUtcWYJB/mafhBLRH8F0Nlj4J8WDu2L51MOJTqeYhZLAoGAWffN
|
||||
icocAbJPlo22sdoa4+/+W5yBF8GAJMDRJtZ+9H1t6SLpQHYRkMIBSETkXUTjZvX9
|
||||
Ocs9iIQkNW9pO/mTdO+VBfCo71JUfknR02xR+6m5gYjlws/ZeYlssXGN2/hbhNiw
|
||||
QOcW7Vv6olFJK6Iy/oz0t6wPO3kpnN3Zogi0paECgYEAwo44M1DdYCtV0snhmYM9
|
||||
5u0mPfYt5P2SVLXyUbr+vFTfrTL/WKnXIJgbsnj3Gvf+GIZv9tKcXhSNmEHQCYX4
|
||||
X3w9iTPddCHuvZ1fpufi2TyArJh0OkoNtLXJHTKrHjf2N+61AQzFiv5WieJrdE+H
|
||||
qr32PTUuVGPyO9LyTY4/RL0=
|
||||
-----END PRIVATE KEY-----
|
||||
";
|
||||
|
|
@ -9,6 +9,7 @@ repository.workspace = true
|
|||
default = []
|
||||
aws = ["dep:litellm-secrets-aws"]
|
||||
google = ["dep:litellm-secrets-google"]
|
||||
hashicorp = ["dep:litellm-secrets-hashicorp"]
|
||||
azure = ["dep:litellm-secrets-azure"]
|
||||
cyberark = ["dep:litellm-secrets-cyberark"]
|
||||
|
||||
|
|
@ -16,6 +17,7 @@ cyberark = ["dep:litellm-secrets-cyberark"]
|
|||
litellm-secrets-types.workspace = true
|
||||
litellm-secrets-aws = { workspace = true, optional = true }
|
||||
litellm-secrets-google = { workspace = true, optional = true }
|
||||
litellm-secrets-hashicorp = { workspace = true, optional = true }
|
||||
litellm-secrets-azure = { workspace = true, optional = true }
|
||||
litellm-secrets-cyberark = { workspace = true, optional = true }
|
||||
litellm-core-utils.workspace = true
|
||||
|
|
|
|||
|
|
@ -9,3 +9,5 @@ Backend failures propagate by default. To allow fallback during a backend failur
|
|||
`get_secret` preserves value types. `get_secret_str` accepts a string default and rejects boolean or JSON values with `Error::TypeMismatch`. `get_secret_bool` accepts a boolean default and converts strings containing `true` or `false`, ignoring surrounding whitespace and ASCII case. Other strings and JSON values produce `Error::TypeMismatch`. Conversion failures never activate fallback or replace a found value with the default
|
||||
|
||||
Provider payloads remain strings unless explicitly selecting a field from an AWS primary JSON secret. Google caches only successfully decoded string payloads, so reads have identical values and types before and after caching. Confirmed absence and failed reads are not cached. AWS resource-not-found responses and Google HTTP 404 responses indicate absence. Other provider errors remain errors, and successful responses without the required payload are malformed responses rather than missing secrets
|
||||
|
||||
The HashiCorp Vault backend is enabled with the `hashicorp` feature and reads KV v2 values from `HCP_VAULT_*` environment variables. It supports static tokens, AppRole authentication, and TLS certificate authentication
|
||||
|
|
|
|||
|
|
@ -30,6 +30,9 @@ pub enum Error {
|
|||
#[cfg(feature = "google")]
|
||||
#[error(transparent)]
|
||||
Google(#[from] litellm_secrets_google::Error),
|
||||
#[cfg(feature = "hashicorp")]
|
||||
#[error(transparent)]
|
||||
Hashicorp(#[from] litellm_secrets_hashicorp::Error),
|
||||
#[cfg(feature = "azure")]
|
||||
#[error(transparent)]
|
||||
Azure(#[from] litellm_secrets_azure::Error),
|
||||
|
|
|
|||
|
|
@ -13,6 +13,8 @@ pub enum SecretManager {
|
|||
GoogleKms(crate::google::GoogleKms),
|
||||
#[cfg(feature = "google")]
|
||||
GoogleSecretManager(crate::google::GoogleSecretManager),
|
||||
#[cfg(feature = "hashicorp")]
|
||||
HashicorpVault(crate::hashicorp::HashicorpVault),
|
||||
#[cfg(feature = "azure")]
|
||||
AzureKeyVault(crate::azure::AzureKeyVault),
|
||||
#[cfg(feature = "cyberark")]
|
||||
|
|
@ -31,6 +33,8 @@ impl SecretManager {
|
|||
Self::GoogleKms(_) => KeyManagementSystem::GoogleKms,
|
||||
#[cfg(feature = "google")]
|
||||
Self::GoogleSecretManager(_) => KeyManagementSystem::GoogleSecretManager,
|
||||
#[cfg(feature = "hashicorp")]
|
||||
Self::HashicorpVault(_) => KeyManagementSystem::HashicorpVault,
|
||||
#[cfg(feature = "azure")]
|
||||
Self::AzureKeyVault(_) => KeyManagementSystem::AzureKeyVault,
|
||||
#[cfg(feature = "cyberark")]
|
||||
|
|
@ -86,6 +90,12 @@ pub async fn get_secret_from_manager(
|
|||
.get_secret_from_google_secret_manager(secret_name)
|
||||
.await
|
||||
.map_err(Error::from),
|
||||
#[cfg(feature = "hashicorp")]
|
||||
SecretManager::HashicorpVault(client) => client
|
||||
.async_read_secret(secret_name)
|
||||
.await
|
||||
.map(|value| value.map(Secret::String))
|
||||
.map_err(Error::from),
|
||||
#[cfg(feature = "azure")]
|
||||
SecretManager::AzureKeyVault(client) => client
|
||||
.get_secret_from_azure_key_vault(secret_name)
|
||||
|
|
|
|||
|
|
@ -23,3 +23,5 @@ pub use litellm_secrets_azure as azure;
|
|||
pub use litellm_secrets_cyberark as cyberark;
|
||||
#[cfg(feature = "google")]
|
||||
pub use litellm_secrets_google as google;
|
||||
#[cfg(feature = "hashicorp")]
|
||||
pub use litellm_secrets_hashicorp as hashicorp;
|
||||
|
|
|
|||
|
|
@ -105,6 +105,148 @@ async fn google_handler_requires_canonical_base64_and_preserves_plaintext_whites
|
|||
Err(Error::MissingCiphertext)
|
||||
));
|
||||
}
|
||||
#[cfg(feature = "hashicorp")]
|
||||
#[tokio::test]
|
||||
async fn hashicorp_handler_resolves_found_missing_and_failed_values() {
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets::{
|
||||
Error, FailurePolicy, KeyManagementSettings, SecretManager, SecretManagerState,
|
||||
SecretResolver, hashicorp::HashicorpVault, hashicorp::HashicorpVaultConfig,
|
||||
};
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{method, path},
|
||||
};
|
||||
|
||||
let found_server = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/v1/secret/data/KEY"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
|
||||
"data": {
|
||||
"data": {"key": "remote"},
|
||||
"metadata": {
|
||||
"created_time": "",
|
||||
"deletion_time": "",
|
||||
"custom_metadata": null,
|
||||
"destroyed": false,
|
||||
"version": 1
|
||||
}
|
||||
},
|
||||
"lease_id": "",
|
||||
"lease_duration": 0,
|
||||
"renewable": false,
|
||||
"request_id": "",
|
||||
"warnings": null,
|
||||
"wrap_info": null
|
||||
})))
|
||||
.mount(&found_server)
|
||||
.await;
|
||||
let found_environment: Arc<dyn Lookup + Send + Sync> = Arc::new({
|
||||
let address = found_server.uri();
|
||||
move |name: &str| match name {
|
||||
"HCP_VAULT_ADDR" => Some(address.clone()),
|
||||
"HCP_VAULT_TOKEN" => Some("token".into()),
|
||||
_ => None,
|
||||
}
|
||||
});
|
||||
let found_config = HashicorpVaultConfig::from_environment(found_environment.as_ref()).unwrap();
|
||||
let found_manager = HashicorpVault::from_config(found_config, true).unwrap();
|
||||
let found_resolver = SecretResolver::new(
|
||||
Arc::new(SecretManagerState::new(
|
||||
SecretManager::HashicorpVault(found_manager),
|
||||
KeyManagementSettings {
|
||||
hosted_keys: Some(vec!["KEY".into()]),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
Arc::new(|_: &str| None),
|
||||
litellm_secrets::OidcResolver::default(),
|
||||
);
|
||||
assert_eq!(
|
||||
found_resolver
|
||||
.get_secret_str("KEY", None)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"remote"
|
||||
);
|
||||
|
||||
let missing_server = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(404).set_body_json(serde_json::json!({"errors": ["missing"]})),
|
||||
)
|
||||
.mount(&missing_server)
|
||||
.await;
|
||||
let missing_environment: Arc<dyn Lookup + Send + Sync> = Arc::new({
|
||||
let address = missing_server.uri();
|
||||
move |name: &str| match name {
|
||||
"HCP_VAULT_ADDR" => Some(address.clone()),
|
||||
"HCP_VAULT_TOKEN" => Some("token".into()),
|
||||
_ => None,
|
||||
}
|
||||
});
|
||||
let missing_config =
|
||||
HashicorpVaultConfig::from_environment(missing_environment.as_ref()).unwrap();
|
||||
let missing_manager = HashicorpVault::from_config(missing_config, true).unwrap();
|
||||
let missing_state = SecretManagerState::new(
|
||||
SecretManager::HashicorpVault(missing_manager),
|
||||
KeyManagementSettings {
|
||||
hosted_keys: Some(vec!["KEY".into()]),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
let missing = litellm_secrets::get_secret_from_manager(
|
||||
missing_state.backend().unwrap(),
|
||||
"KEY",
|
||||
missing_state.settings().unwrap(),
|
||||
&|_: &str| None,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(missing.is_none());
|
||||
|
||||
let failed_server = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(500).set_body_json(serde_json::json!({"errors": ["failed"]})),
|
||||
)
|
||||
.mount(&failed_server)
|
||||
.await;
|
||||
let failed_environment: Arc<dyn Lookup + Send + Sync> = Arc::new({
|
||||
let address = failed_server.uri();
|
||||
move |name: &str| match name {
|
||||
"HCP_VAULT_ADDR" => Some(address.clone()),
|
||||
"HCP_VAULT_TOKEN" => Some("token".into()),
|
||||
_ => None,
|
||||
}
|
||||
});
|
||||
let failed_config =
|
||||
HashicorpVaultConfig::from_environment(failed_environment.as_ref()).unwrap();
|
||||
let failed_manager = HashicorpVault::from_config(failed_config, true).unwrap();
|
||||
let failed_state = SecretManagerState::new(
|
||||
SecretManager::HashicorpVault(failed_manager),
|
||||
KeyManagementSettings {
|
||||
hosted_keys: Some(vec!["KEY".into()]),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
let failed_resolver = SecretResolver::new(
|
||||
Arc::new(failed_state),
|
||||
Arc::new(|_: &str| None),
|
||||
litellm_secrets::OidcResolver::default(),
|
||||
)
|
||||
.with_failure_policy(FailurePolicy::Propagate);
|
||||
assert!(matches!(
|
||||
failed_resolver.get_secret_str("KEY", None).await,
|
||||
Err(Error::Hashicorp(
|
||||
litellm_secrets::hashicorp::Error::Status { status: 500 }
|
||||
))
|
||||
));
|
||||
}
|
||||
|
||||
#[cfg(feature = "azure")]
|
||||
#[tokio::test]
|
||||
|
|
|
|||
|
|
@ -1860,6 +1860,12 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [
|
|||
"max_ui_session_budget",
|
||||
"budget_rollover",
|
||||
"mcp_tool_search",
|
||||
"turn_off_message_logging",
|
||||
"datadog_params",
|
||||
"datadog_llm_observability_params",
|
||||
"newrelic_params",
|
||||
"pointfive_params",
|
||||
"aws_sqs_callback_params",
|
||||
]
|
||||
SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"]
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))
|
||||
|
|
|
|||
|
|
@ -118,6 +118,7 @@ class XAIChatConfig(OpenAIGPTConfig):
|
|||
base_openai_params: Final = [
|
||||
"logit_bias",
|
||||
"logprobs",
|
||||
"max_completion_tokens",
|
||||
"max_tokens",
|
||||
"n",
|
||||
"parallel_tool_calls",
|
||||
|
|
|
|||
|
|
@ -114,13 +114,22 @@ class _CacheTestHandle:
|
|||
@staticmethod
|
||||
def disk(directory: str) -> _CacheTestHandle: ...
|
||||
@staticmethod
|
||||
def gcs(
|
||||
bucket_name: str,
|
||||
*,
|
||||
gcs_path: str | None = None,
|
||||
path_service_account: str | None = None,
|
||||
endpoint: str | None = None,
|
||||
token: str | None = None,
|
||||
) -> _CacheTestHandle: ...
|
||||
@staticmethod
|
||||
def azure_blob(account_url: str, container: str) -> _CacheTestHandle: ...
|
||||
@staticmethod
|
||||
def redis_semantic(backend: object) -> _CacheTestHandle: ...
|
||||
@property
|
||||
def backend(
|
||||
self,
|
||||
) -> Literal["memory", "redis", "disk", "azure-blob", "redis_semantic"]: ...
|
||||
) -> Literal["memory", "redis", "gcs", "disk", "azure-blob", "redis_semantic"]: ...
|
||||
def _bind_facade(self, facade: object) -> None: ...
|
||||
|
||||
@final
|
||||
|
|
|
|||
|
|
@ -90,6 +90,16 @@ class TestXAIReasoningTokenFolding:
|
|||
assert response.usage.total_tokens == 999
|
||||
|
||||
|
||||
def test_max_completion_tokens_is_accepted_and_mapped_to_max_tokens() -> None:
|
||||
optional_params = litellm.get_optional_params(
|
||||
model="grok-4.20",
|
||||
custom_llm_provider="xai",
|
||||
max_completion_tokens=64,
|
||||
)
|
||||
assert optional_params["max_tokens"] == 64, optional_params
|
||||
assert "max_completion_tokens" not in optional_params, optional_params
|
||||
|
||||
|
||||
class TestXAIParallelToolCalls:
|
||||
"""Test suite for XAI parallel tool calls functionality."""
|
||||
|
||||
|
|
|
|||
|
|
@ -11766,6 +11766,66 @@ def test_prompt_caching_settings_propagate_on_config_reload(monkeypatch, field_n
|
|||
assert getattr(litellm, field_name) == db_value
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_stored_datadog_redaction_settings_apply_before_logger_init(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A DB-only litellm_settings row that pairs success_callback: ["datadog"] with
|
||||
datadog_params.turn_off_message_logging: true must build the DataDogLogger redacted, the
|
||||
same as the identical block in YAML. Regression for the redaction keys being absent from
|
||||
the safe-override allowlist while the callback half of the row was honoured."""
|
||||
import litellm.proxy.proxy_server as ps
|
||||
from litellm.integrations.datadog.datadog import DataDogLogger
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
|
||||
monkeypatch.setenv("DD_API_KEY", "test-key")
|
||||
monkeypatch.setenv("DD_SITE", "us5.datadoghq.com")
|
||||
monkeypatch.setattr(litellm, "datadog_params", None)
|
||||
monkeypatch.setattr(litellm, "turn_off_message_logging", False)
|
||||
monkeypatch.setattr(litellm, "success_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm_logging, "_in_memory_loggers", [])
|
||||
|
||||
db_row = {
|
||||
"success_callback": ["datadog"],
|
||||
"datadog_params": {"turn_off_message_logging": True},
|
||||
"turn_off_message_logging": True,
|
||||
}
|
||||
pc = ps.ProxyConfig()
|
||||
pc._apply_litellm_settings_db_values(pc._prepared_db_settings_values("litellm_settings", db_row))
|
||||
pc._add_callbacks_from_db_config({"litellm_settings": db_row})
|
||||
|
||||
datadog_loggers = [cb for cb in litellm.success_callback if isinstance(cb, DataDogLogger)]
|
||||
assert len(datadog_loggers) == 1
|
||||
assert datadog_loggers[0].turn_off_message_logging is True
|
||||
assert litellm.turn_off_message_logging is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"field_name",
|
||||
[
|
||||
"datadog_params",
|
||||
"datadog_llm_observability_params",
|
||||
"newrelic_params",
|
||||
"pointfive_params",
|
||||
"aws_sqs_callback_params",
|
||||
],
|
||||
)
|
||||
def test_db_stored_callback_params_propagate_to_litellm_module(monkeypatch: pytest.MonkeyPatch, field_name: str):
|
||||
"""Every callback init params block stored in the DB litellm_settings row must land on the
|
||||
litellm module before the matching logger is built, so the DB row behaves like YAML."""
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
monkeypatch.setattr(litellm, field_name, None)
|
||||
db_value = {"turn_off_message_logging": True}
|
||||
|
||||
pc = ps.ProxyConfig()
|
||||
pc._apply_litellm_settings_db_values(pc._prepared_db_settings_values("litellm_settings", {field_name: db_value}))
|
||||
|
||||
assert getattr(litellm, field_name) == db_value
|
||||
|
||||
|
||||
def test_get_config_list_marks_untouched_prompt_caching_flag_as_not_set(monkeypatch):
|
||||
"""The flag defaults to False rather than None, so a plain 'is not None' check would
|
||||
report the default as 'In Config' and imply an admin had set it."""
|
||||
|
|
|
|||
|
|
@ -0,0 +1,85 @@
|
|||
[
|
||||
{
|
||||
"name": "defaults",
|
||||
"env": {
|
||||
"HCP_VAULT_TOKEN": "token"
|
||||
},
|
||||
"secret_name": "OPENAI_API_KEY",
|
||||
"expected_secret_url": "http://127.0.0.1:8200/v1/secret/data/OPENAI_API_KEY",
|
||||
"expected_login_url": null,
|
||||
"expected_login_namespace": null,
|
||||
"expected_secret_namespace": null
|
||||
},
|
||||
{
|
||||
"name": "global_namespace",
|
||||
"env": {
|
||||
"HCP_VAULT_ADDR": "http://vault.test:8200",
|
||||
"HCP_VAULT_TOKEN": "token",
|
||||
"HCP_VAULT_NAMESPACE": "admin"
|
||||
},
|
||||
"secret_name": "OPENAI_API_KEY",
|
||||
"expected_secret_url": "http://vault.test:8200/v1/admin/secret/data/OPENAI_API_KEY",
|
||||
"expected_login_url": null,
|
||||
"expected_login_namespace": "admin",
|
||||
"expected_secret_namespace": "admin"
|
||||
},
|
||||
{
|
||||
"name": "namespace_overrides",
|
||||
"env": {
|
||||
"HCP_VAULT_ADDR": "http://vault.test:8200",
|
||||
"HCP_VAULT_TOKEN": "token",
|
||||
"HCP_VAULT_NAMESPACE": "admin",
|
||||
"HCP_VAULT_LOGIN_NAMESPACE": "root",
|
||||
"HCP_VAULT_SECRET_NAMESPACE": "teams/team-a"
|
||||
},
|
||||
"secret_name": "OPENAI_API_KEY",
|
||||
"expected_secret_url": "http://vault.test:8200/v1/teams/team-a/secret/data/OPENAI_API_KEY",
|
||||
"expected_login_url": null,
|
||||
"expected_login_namespace": "root",
|
||||
"expected_secret_namespace": "teams/team-a"
|
||||
},
|
||||
{
|
||||
"name": "custom_mount_and_prefix",
|
||||
"env": {
|
||||
"HCP_VAULT_ADDR": "http://vault.test:8200",
|
||||
"HCP_VAULT_TOKEN": "token",
|
||||
"HCP_VAULT_MOUNT_NAME": " /kv-prod/ ",
|
||||
"HCP_VAULT_PATH_PREFIX": " /virtual-keys/ "
|
||||
},
|
||||
"secret_name": "DB_PASSWORD",
|
||||
"expected_secret_url": "http://vault.test:8200/v1/kv-prod/data/virtual-keys/DB_PASSWORD",
|
||||
"expected_login_url": null,
|
||||
"expected_login_namespace": null,
|
||||
"expected_secret_namespace": null
|
||||
},
|
||||
{
|
||||
"name": "approle_custom_mount",
|
||||
"env": {
|
||||
"HCP_VAULT_ADDR": "http://vault.test:8200",
|
||||
"HCP_VAULT_APPROLE_ROLE_ID": "role-id",
|
||||
"HCP_VAULT_APPROLE_SECRET_ID": "secret-id",
|
||||
"HCP_VAULT_APPROLE_MOUNT_PATH": "custom-approle",
|
||||
"HCP_VAULT_NAMESPACE": "admin"
|
||||
},
|
||||
"secret_name": "OPENAI_API_KEY",
|
||||
"expected_secret_url": "http://vault.test:8200/v1/admin/secret/data/OPENAI_API_KEY",
|
||||
"expected_login_url": "http://vault.test:8200/v1/auth/custom-approle/login",
|
||||
"expected_login_namespace": "admin",
|
||||
"expected_secret_namespace": "admin"
|
||||
},
|
||||
{
|
||||
"name": "tls_cert",
|
||||
"env": {
|
||||
"HCP_VAULT_ADDR": "http://vault.test:8200",
|
||||
"HCP_VAULT_CLIENT_CERT": "/tmp/client.crt",
|
||||
"HCP_VAULT_CLIENT_KEY": "/tmp/client.key",
|
||||
"HCP_VAULT_CERT_ROLE": "vault-role",
|
||||
"HCP_VAULT_NAMESPACE": "admin"
|
||||
},
|
||||
"secret_name": "OPENAI_API_KEY",
|
||||
"expected_secret_url": "http://vault.test:8200/v1/admin/secret/data/OPENAI_API_KEY",
|
||||
"expected_login_url": "http://vault.test:8200/v1/auth/cert/login",
|
||||
"expected_login_namespace": "admin",
|
||||
"expected_secret_namespace": "admin"
|
||||
}
|
||||
]
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
import datetime
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
|
@ -18,6 +19,23 @@ LOGIN_RESPONSE: Final = {"auth": {"client_token": "hvs.login-token", "lease_dura
|
|||
SECRET_RESPONSE: Final = {"data": {"data": {"key": "sk-from-vault", "password": "pw-from-vault"}}}
|
||||
|
||||
NAMESPACE_ENV_VARS: Final = ("HCP_VAULT_NAMESPACE", "HCP_VAULT_LOGIN_NAMESPACE", "HCP_VAULT_SECRET_NAMESPACE")
|
||||
PARITY_ENV_VARS: Final = (
|
||||
"HCP_VAULT_ADDR",
|
||||
"HCP_VAULT_TOKEN",
|
||||
"HCP_VAULT_NAMESPACE",
|
||||
"HCP_VAULT_LOGIN_NAMESPACE",
|
||||
"HCP_VAULT_SECRET_NAMESPACE",
|
||||
"HCP_VAULT_MOUNT_NAME",
|
||||
"HCP_VAULT_PATH_PREFIX",
|
||||
"HCP_VAULT_APPROLE_ROLE_ID",
|
||||
"HCP_VAULT_APPROLE_SECRET_ID",
|
||||
"HCP_VAULT_APPROLE_MOUNT_PATH",
|
||||
"HCP_VAULT_CLIENT_CERT",
|
||||
"HCP_VAULT_CLIENT_KEY",
|
||||
"HCP_VAULT_CERT_ROLE",
|
||||
"HCP_VAULT_REFRESH_INTERVAL",
|
||||
"SECRET_MANAGER_REFRESH_INTERVAL",
|
||||
)
|
||||
|
||||
|
||||
def _build_manager(monkeypatch: pytest.MonkeyPatch, env: Mapping[str, str]) -> HashicorpSecretManager:
|
||||
|
|
@ -236,3 +254,35 @@ def test_tls_login_uses_login_namespace(monkeypatch: pytest.MonkeyPatch, tmp_pat
|
|||
|
||||
assert manager._auth_via_tls_cert() == "hvs.login-token"
|
||||
assert login_route.calls.last.request.headers["X-Vault-Namespace"] == "root"
|
||||
|
||||
|
||||
with Path(__file__).with_name("hashicorp_vault_parity.json").open() as parity_file:
|
||||
PARITY_CASES: Final = json.load(parity_file)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("case", PARITY_CASES, ids=lambda case: case["name"])
|
||||
def test_configuration_matches_native_parity_fixture(
|
||||
monkeypatch: pytest.MonkeyPatch, case: Mapping[str, object]
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True)
|
||||
for name in PARITY_ENV_VARS:
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
for name, value in case["env"].items():
|
||||
monkeypatch.setenv(name, value)
|
||||
|
||||
manager: Final = HashicorpSecretManager()
|
||||
env: Final = case["env"]
|
||||
expected_login_url: Final = case["expected_login_url"]
|
||||
if env.get("HCP_VAULT_APPROLE_ROLE_ID") and env.get("HCP_VAULT_APPROLE_SECRET_ID"):
|
||||
login_url: str | None = (
|
||||
f"{manager.vault_addr}/v1/auth/{manager.approle_mount_path}/login"
|
||||
)
|
||||
elif env.get("HCP_VAULT_CLIENT_CERT") and env.get("HCP_VAULT_CLIENT_KEY"):
|
||||
login_url = f"{manager.vault_addr}/v1/auth/cert/login"
|
||||
else:
|
||||
login_url = None
|
||||
|
||||
assert manager.get_url(case["secret_name"]) == case["expected_secret_url"]
|
||||
assert manager.vault_login_namespace == case["expected_login_namespace"]
|
||||
assert manager.vault_secret_namespace == case["expected_secret_namespace"]
|
||||
assert login_url == expected_login_url
|
||||
|
|
|
|||
152
tests/test_litellm_rust/support/fake_gcs.py
Normal file
152
tests/test_litellm_rust/support/fake_gcs.py
Normal file
|
|
@ -0,0 +1,152 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from functools import partial
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from socket import socket
|
||||
from types import MappingProxyType
|
||||
from typing import Final, cast
|
||||
from urllib.parse import unquote, urlsplit
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RecordedRequest:
|
||||
method: str
|
||||
path: str
|
||||
query: str
|
||||
headers: Mapping[str, str]
|
||||
body: bytes
|
||||
|
||||
|
||||
class _FakeGcsHandler(BaseHTTPRequestHandler):
|
||||
def __init__(
|
||||
self,
|
||||
request: socket | tuple[bytes, socket],
|
||||
client_address: tuple[str, int],
|
||||
server: ThreadingHTTPServer,
|
||||
*,
|
||||
fake: FakeGcs,
|
||||
) -> None:
|
||||
self._fake: Final = fake
|
||||
super().__init__(request, client_address, server)
|
||||
|
||||
def _handle(self) -> None:
|
||||
parsed: Final = urlsplit(self.path)
|
||||
content_length: Final = int(self.headers.get("Content-Length", "0"))
|
||||
body: Final = self.rfile.read(content_length) if content_length else b""
|
||||
headers: Final = MappingProxyType(
|
||||
{name.title(): value for name, value in self.headers.items()}
|
||||
)
|
||||
self._fake.record(
|
||||
RecordedRequest(
|
||||
method=self.command,
|
||||
path=parsed.path,
|
||||
query=parsed.query,
|
||||
headers=headers,
|
||||
body=body,
|
||||
)
|
||||
)
|
||||
if self.headers.get("Authorization") != f"Bearer {self._fake.token}":
|
||||
self._send_json(401, {"error": "unauthorized"})
|
||||
return
|
||||
|
||||
upload_prefix: Final = "/upload/storage/v1/b/"
|
||||
download_prefix: Final = "/storage/v1/b/"
|
||||
if parsed.path.startswith(upload_prefix) and parsed.path.endswith("/o"):
|
||||
self._upload(parsed.path[len(upload_prefix) : -2], parsed.query, body)
|
||||
return
|
||||
if parsed.path.startswith(download_prefix):
|
||||
self._download(parsed.path[len(download_prefix) :], parsed.query)
|
||||
return
|
||||
self._send_json(404, {"error": "not found"})
|
||||
|
||||
def _upload(self, path: str, query: str, body: bytes) -> None:
|
||||
values: Final = {
|
||||
unquote(pair.partition("=")[0]): unquote(pair.partition("=")[2])
|
||||
for pair in query.split("&")
|
||||
if pair
|
||||
}
|
||||
if not path or values.get("uploadType") != "media" or "name" not in values:
|
||||
self._send_json(404, {"error": "not found"})
|
||||
return
|
||||
self._fake.put_object(path, values["name"], body)
|
||||
self._send_json(200, {"name": values["name"], "bucket": path})
|
||||
|
||||
def _download(self, path: str, query: str) -> None:
|
||||
bucket, separator, encoded_name = path.partition("/o/")
|
||||
if not separator or query != "alt=media":
|
||||
self._send_json(404, {"error": "not found"})
|
||||
return
|
||||
name: Final = unquote(encoded_name)
|
||||
if name.endswith("/server-error") or name == "server-error":
|
||||
self._send_json(500, {"error": "server error"})
|
||||
return
|
||||
body: Final = self._fake.get_object(bucket, name)
|
||||
if body is None:
|
||||
self._send_json(404, {"error": "not found"})
|
||||
return
|
||||
self._send(200, body, "application/octet-stream")
|
||||
|
||||
def _send_json(self, status: int, value: object) -> None:
|
||||
payload: Final = json.dumps(value).encode()
|
||||
self._send(status, payload, "application/json")
|
||||
|
||||
def _send(self, status: int, body: bytes, content_type: str) -> None:
|
||||
self.send_response(status)
|
||||
self.send_header("Content-Type", content_type)
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
def log_message(self, format: str, *args: object) -> None:
|
||||
pass
|
||||
|
||||
do_GET = _handle
|
||||
do_POST = _handle
|
||||
|
||||
|
||||
class FakeGcs:
|
||||
def __init__(self) -> None:
|
||||
self._objects: dict[tuple[str, str], bytes] = {} # mutable-ok: fake object store
|
||||
self._requests: list[RecordedRequest] = [] # mutable-ok: recorded request history
|
||||
self._server = ThreadingHTTPServer(
|
||||
("127.0.0.1", 0),
|
||||
partial(_FakeGcsHandler, fake=self),
|
||||
)
|
||||
self._worker = threading.Thread(target=self._server.serve_forever, daemon=True)
|
||||
self._worker.start()
|
||||
self.token: Final = "test-token"
|
||||
|
||||
@property
|
||||
def url(self) -> str:
|
||||
address: Final = cast(tuple[str, int], self._server.server_address)
|
||||
host, port = address
|
||||
return f"http://{host}:{port}"
|
||||
|
||||
@property
|
||||
def objects(self) -> Mapping[tuple[str, str], bytes]:
|
||||
return MappingProxyType(self._objects)
|
||||
|
||||
@property
|
||||
def requests(self) -> tuple[RecordedRequest, ...]:
|
||||
return tuple(self._requests)
|
||||
|
||||
def put(self, bucket: str, name: str, body: bytes) -> None:
|
||||
self.put_object(bucket, name, body)
|
||||
|
||||
def close(self) -> None:
|
||||
self._server.shutdown()
|
||||
self._server.server_close()
|
||||
self._worker.join(timeout=5)
|
||||
|
||||
def record(self, request: RecordedRequest) -> None:
|
||||
self._requests.append(request)
|
||||
|
||||
def put_object(self, bucket: str, name: str, body: bytes) -> None:
|
||||
self._objects[(bucket, name)] = body
|
||||
|
||||
def get_object(self, bucket: str, name: str) -> bytes | None:
|
||||
return self._objects.get((bucket, name))
|
||||
|
|
@ -26,6 +26,7 @@ from azure.storage.blob import ContainerClient
|
|||
import litellm
|
||||
from litellm.caching.azure_blob_cache import AzureBlobCache
|
||||
from litellm.caching.caching import Cache, disable_cache, enable_cache, update_cache
|
||||
from litellm.caching.gcs_cache import GCSCache
|
||||
from litellm.caching.disk_cache import DiskCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.caching.redis_cluster_cache import RedisClusterCache
|
||||
|
|
@ -34,6 +35,7 @@ from litellm.rust_bridge import _native
|
|||
from litellm.types.caching import LiteLLMCacheType
|
||||
from litellm.types.llms.custom_llm import CustomLLMItem
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from tests.test_litellm_rust.support.fake_gcs import FakeGcs
|
||||
from tests.test_litellm_rust.support.isolation import rebound
|
||||
|
||||
_CacheTestHandle: Final = _native._CacheTestHandle # pyright: ignore[reportPrivateUsage] # test-only handle has no public module name
|
||||
|
|
@ -44,6 +46,7 @@ pytestmark: Final = pytest.mark.requires_rust_extension
|
|||
|
||||
class CacheLookup(Protocol):
|
||||
def get_cache(self, **kwargs: object) -> object: ...
|
||||
def flush_cache(self) -> object: ...
|
||||
|
||||
|
||||
def request(key: str = "key") -> dict[str, object]:
|
||||
|
|
@ -63,6 +66,15 @@ def redis_url() -> Generator[str]:
|
|||
worker.join(timeout=5)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_gcs() -> Generator[FakeGcs]:
|
||||
server: Final = FakeGcs()
|
||||
try:
|
||||
yield server
|
||||
finally:
|
||||
server.close()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def azure_blob_facade() -> Generator[Cache]:
|
||||
account_url: Final = os.environ.get("AZURE_BLOB_CACHE_ACCOUNT_URL")
|
||||
|
|
@ -640,6 +652,222 @@ async def test_disk_native_batch_lookup_and_store_report_partial_hits(tmp_path:
|
|||
}
|
||||
|
||||
|
||||
async def test_gcs_reads_python_entries_and_writes_python_compatible_objects(
|
||||
fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
|
||||
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
|
||||
response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}, "flag": True, "empty": None}
|
||||
fake_gcs.put(
|
||||
"bucket",
|
||||
"cache/sync",
|
||||
json.dumps({"timestamp": time.time(), "response": json.dumps(response)}).encode(),
|
||||
)
|
||||
fake_gcs.put("bucket", "cache/async", json.dumps({"timestamp": time.time(), "response": response}).encode())
|
||||
fake_gcs.put("bucket", "cache/raw", json.dumps(response).encode())
|
||||
fake_gcs.put("bucket", "cache/invalid", b"not a cache entry")
|
||||
binding: Final = _native._CacheTestResolver(
|
||||
SimpleNamespace(
|
||||
cache=_native._CacheTestHandle.gcs(
|
||||
"bucket",
|
||||
gcs_path="cache",
|
||||
endpoint=fake_gcs.url,
|
||||
token=fake_gcs.token,
|
||||
)
|
||||
)
|
||||
).resolve()
|
||||
|
||||
assert binding.lookup(request("sync")) == response
|
||||
assert await binding.async_lookup(request("async")) == response
|
||||
assert binding.lookup(request("raw")) == response
|
||||
assert await binding.async_lookup(request("invalid")) is None
|
||||
assert binding.lookup(request("missing")) is None
|
||||
|
||||
await binding.async_store({**request("native"), "ttl_seconds": 12.0}, response)
|
||||
stored: Final = fake_gcs.objects[("bucket", "cache/native")]
|
||||
stored_value: Final = cast(dict[str, object], json.loads(stored))
|
||||
assert stored_value["response"] == response
|
||||
assert isinstance(stored_value["timestamp"], float)
|
||||
upload: Final = next(item for item in fake_gcs.requests if item.method == "POST")
|
||||
assert upload.path == "/upload/storage/v1/b/bucket/o"
|
||||
assert upload.query == "uploadType=media&name=cache%2Fnative"
|
||||
assert upload.headers["Authorization"] == f"Bearer {fake_gcs.token}"
|
||||
assert upload.headers["Content-Type"] == "application/json"
|
||||
upload_text: Final = f"{upload.path}?{upload.query}{upload.headers}"
|
||||
assert "ttl" not in upload_text.lower()
|
||||
assert "expiry" not in upload_text.lower()
|
||||
download: Final = next(item for item in fake_gcs.requests if item.path.endswith("/cache%2Fsync"))
|
||||
assert download.path == "/storage/v1/b/bucket/o/cache%2Fsync"
|
||||
assert download.query == "alt=media"
|
||||
|
||||
binding.store(request("sync2"), response)
|
||||
assert binding.lookup(request("sync2")) == response
|
||||
assert GCSCache(bucket_name="bucket", gcs_path="cache").key_prefix == "cache/"
|
||||
assert GCSCache(bucket_name="bucket", gcs_path="cache/").key_prefix == "cache/"
|
||||
assert GCSCache(bucket_name="bucket").key_prefix == ""
|
||||
|
||||
|
||||
async def test_gcs_batch_lookup_preserves_order_and_treats_malformed_entries_as_misses(fake_gcs: FakeGcs) -> None:
|
||||
fake_gcs.put("bucket", "cache/hit", json.dumps({"timestamp": time.time(), "response": {"value": 1}}).encode())
|
||||
fake_gcs.put("bucket", "cache/invalid", b"not a cache entry")
|
||||
binding: Final = _native._CacheTestResolver(
|
||||
SimpleNamespace(
|
||||
cache=_native._CacheTestHandle.gcs(
|
||||
"bucket",
|
||||
gcs_path="cache",
|
||||
endpoint=fake_gcs.url,
|
||||
token=fake_gcs.token,
|
||||
)
|
||||
)
|
||||
).resolve()
|
||||
requests: Final = [request("hit"), request("missing"), request("invalid")]
|
||||
expected: Final = {"values": [{"value": 1}, None, None], "missing_indices": [1, 2]}
|
||||
|
||||
assert await binding.async_lookup_batch(requests) == expected
|
||||
assert binding.lookup_batch(requests) == expected
|
||||
await binding.async_store_batch([request("first"), request("second")], [{"value": 1}, {"value": 2}])
|
||||
assert ("bucket", "cache/first") in fake_gcs.objects
|
||||
assert ("bucket", "cache/second") in fake_gcs.objects
|
||||
|
||||
|
||||
async def test_gcs_facade_binds_only_exact_matching_configuration(
|
||||
fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
|
||||
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
|
||||
monkeypatch.setenv("GOOGLE_APPLICATION_CREDENTIALS", "/nonexistent")
|
||||
facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/")
|
||||
assert type(facade.cache) is GCSCache
|
||||
|
||||
mismatched_bucket: Final = _native._CacheTestHandle.gcs(
|
||||
"other",
|
||||
gcs_path="cache",
|
||||
endpoint=fake_gcs.url,
|
||||
token=fake_gcs.token,
|
||||
)
|
||||
with pytest.raises(TypeError, match="buckets must match"):
|
||||
mismatched_bucket._bind_facade(facade)
|
||||
mismatched_prefix: Final = _native._CacheTestHandle.gcs(
|
||||
"bucket",
|
||||
gcs_path="x",
|
||||
endpoint=fake_gcs.url,
|
||||
token=fake_gcs.token,
|
||||
)
|
||||
with pytest.raises(TypeError, match="key prefixes must match"):
|
||||
mismatched_prefix._bind_facade(facade)
|
||||
mismatched_credentials: Final = _native._CacheTestHandle.gcs(
|
||||
"bucket",
|
||||
gcs_path="cache",
|
||||
path_service_account="sa.json",
|
||||
endpoint=fake_gcs.url,
|
||||
token=fake_gcs.token,
|
||||
)
|
||||
with pytest.raises(TypeError, match="credentials must match"):
|
||||
mismatched_credentials._bind_facade(facade)
|
||||
with pytest.raises(TypeError, match="types must match"):
|
||||
_native._CacheTestHandle.memory()._bind_facade(facade)
|
||||
|
||||
matching: Final = _native._CacheTestHandle.gcs(
|
||||
"bucket",
|
||||
gcs_path="cache",
|
||||
endpoint=fake_gcs.url,
|
||||
token=fake_gcs.token,
|
||||
)
|
||||
matching._bind_facade(facade)
|
||||
resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade))
|
||||
binding: Final = resolver.resolve()
|
||||
assert binding.kind == "native"
|
||||
await binding.async_store(request("native"), {"value": "native"})
|
||||
assert await binding.async_lookup(request("native")) == {"value": "native"}
|
||||
assert cast(CacheLookup, facade).get_cache(cache_key="native") is None
|
||||
|
||||
with rebound(facade.cache, "bucket_name", "other"):
|
||||
assert resolver.resolve().kind == "python_callback"
|
||||
with rebound(facade.cache, "key_prefix", "x/"):
|
||||
assert resolver.resolve().kind == "python_callback"
|
||||
with rebound(facade.cache, "path_service_account", "sa.json"):
|
||||
assert resolver.resolve().kind == "python_callback"
|
||||
def no_get_cache(*args: object, **kwargs: object) -> None:
|
||||
return None
|
||||
|
||||
with rebound(facade.cache, "get_cache", no_get_cache):
|
||||
assert resolver.resolve().kind == "python_callback"
|
||||
with rebound(facade, "ttl", 12):
|
||||
assert resolver.resolve().kind == "python_callback"
|
||||
|
||||
class CustomGcs(GCSCache):
|
||||
pass
|
||||
|
||||
with rebound(facade, "cache", CustomGcs(bucket_name="bucket", gcs_path="cache/")):
|
||||
assert resolver.resolve().kind == "python_callback"
|
||||
custom_facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/")
|
||||
with rebound(custom_facade, "cache", CustomGcs(bucket_name="bucket", gcs_path="cache/")):
|
||||
with pytest.raises(TypeError, match="types must match"):
|
||||
matching._bind_facade(custom_facade)
|
||||
|
||||
missing_bucket: Final = Cache(type=LiteLLMCacheType.GCS)
|
||||
with pytest.raises(TypeError, match="requires a configured bucket name"):
|
||||
matching._bind_facade(missing_bucket)
|
||||
|
||||
|
||||
async def test_gcs_flush_is_a_no_op_and_ping_is_not_implemented(
|
||||
fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
|
||||
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
|
||||
binding: Final = _native._CacheTestResolver(
|
||||
SimpleNamespace(
|
||||
cache=_native._CacheTestHandle.gcs(
|
||||
"bucket",
|
||||
gcs_path="cache",
|
||||
endpoint=fake_gcs.url,
|
||||
token=fake_gcs.token,
|
||||
)
|
||||
)
|
||||
).resolve()
|
||||
await binding.async_store(request("key"), {"value": "stored"})
|
||||
await binding.async_flush()
|
||||
assert ("bucket", "cache/key") in fake_gcs.objects
|
||||
assert await binding.async_lookup(request("key")) == {"value": "stored"}
|
||||
with pytest.raises(NotImplementedError):
|
||||
await binding.ping()
|
||||
|
||||
facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/")
|
||||
with pytest.raises(AttributeError):
|
||||
await facade.ping()
|
||||
assert cast(CacheLookup, facade.cache).flush_cache() is None
|
||||
|
||||
|
||||
async def test_gcs_unauthorized_and_server_errors_surface_as_runtime_errors(fake_gcs: FakeGcs) -> None:
|
||||
wrong_token: Final = _native._CacheTestResolver(
|
||||
SimpleNamespace(
|
||||
cache=_native._CacheTestHandle.gcs(
|
||||
"bucket",
|
||||
gcs_path="cache",
|
||||
endpoint=fake_gcs.url,
|
||||
token="wrong-token",
|
||||
)
|
||||
)
|
||||
).resolve()
|
||||
with pytest.raises(RuntimeError):
|
||||
wrong_token.lookup(request("missing"))
|
||||
assert not fake_gcs.objects
|
||||
|
||||
binding: Final = _native._CacheTestResolver(
|
||||
SimpleNamespace(
|
||||
cache=_native._CacheTestHandle.gcs(
|
||||
"bucket",
|
||||
gcs_path="cache",
|
||||
endpoint=fake_gcs.url,
|
||||
token=fake_gcs.token,
|
||||
)
|
||||
)
|
||||
).resolve()
|
||||
with pytest.raises(RuntimeError):
|
||||
binding.lookup(request("server-error"))
|
||||
assert binding.lookup(request("missing")) is None
|
||||
|
||||
|
||||
async def test_redis_cluster_facade_serves_multi_slot_batches_and_scoped_flush_natively(
|
||||
cluster_nodes: tuple[tuple[str, int], ...],
|
||||
) -> None:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue