refactor(http): hand out an owned Client and route all providers through the pool (#43245)

* refactor(messages): take the provider client from the injected HTTP pool

The messages route kept its own process-wide reqwest client, so it ignored
ssl_verify, CA bundles, client certs, proxies and every other setting that
litellm-http resolves. The machine now takes the HttpClientPool and the
call's HttpClientConfig, as OCR does, and the bridge passes its shared pool.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* refactor(http): hand out an owned Client and move chat, audio and OIDC onto the pool

HttpClientPool now returns litellm_http::Client, a newtype only crates/http
can build, so every provider client carries the resolved TLS, proxy and
timeout settings. Chat completions and audio transcription drop their
process-wide reqwest clients and take the pool and call config like
messages; their 600s ceiling moves to the request. OidcResolver takes its
client instead of building one, and the bridge hands it the pooled one.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* refactor(secrets): build Google, Azure and CyberArk manager clients from the pool

The native secret managers built bare reqwest clients, so they ignored the
host's TLS and proxy settings. load_native_manager now takes the pool and
the host config and hands each manager a pooled client.

CyberArk's CYBERARK_SSL_VERIFY and CYBERARK_CLIENT_CERT/KEY become an
override on the host config instead of a hand-built client. To express a
certificate and key in separate files, HttpClientConfig::client_certificate
is now a ClientIdentity that is either one PEM or a split pair.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* chore(clippy): only crates/http may build a reqwest client

Fence reqwest::Client, ClientBuilder and the TLS builder methods with
disallowed-types and disallowed-methods so new code takes a
litellm_http::Client from the pool. crates/http is exempt as the one place
clients are built, and testkit as a dev-only installer. Tests move to
litellm_http::Client::plain_for_test or a pooled client.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* fix(secrets-cyberark): keep verifying certificates when the host disables it

Python hands CyberArk its own ssl_verify, which wins over the global
setting, so CYBERARK_SSL_VERIFY unset or true still verifies even when the
host sets ssl_verify false. The pooled client copied the host's Disabled
and would send the API key unverified; fall back to the built-in roots
instead.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* fix(python-bridge): treat a missing litellm package as no host HTTP settings

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

---------

Co-authored-by: Yujong Lee <yujong@berri.ai>
Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-25 18:31:15 -07:00 • committed by GitHub
parent 4adbc13d79
commit 1a58162630
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
102 changed files with 745 additions and 403 deletions

View file

@ -2911,6 +2911,7 @@ dependencies = [
"litellm-cache",
"litellm-cache-response",
"litellm-cache-testing",
"litellm-http",
"reqwest 0.12.28",
"rstest",
"serde_json",
@ -2944,6 +2945,7 @@ dependencies = [
"litellm-auth-types",
"litellm-cache",
"litellm-cache-testing",
"litellm-http",
"percent-encoding",
"reqwest 0.12.28",
"rstest",
@ -2970,6 +2972,7 @@ dependencies = [
"futures-util",
"litellm-cache",
"litellm-cache-testing",
"litellm-http",
"qdrant-client",
"reqwest 0.12.28",
"rstest",
@ -3043,6 +3046,7 @@ dependencies = [
"litellm-auth-aws",
"litellm-cache",
"litellm-cache-testing",
"litellm-http",
"reqwest 0.12.28",
"rstest",
"serde_json",
@ -3207,11 +3211,13 @@ dependencies = [
"http 1.4.2",
"hyper-util",
"litellm-core-utils",
"rcgen",
"reqwest 0.12.28",
"rstest",
"rustls 0.23.42",
"serde",
"serde_json",
"tempfile",
"thiserror 2.0.19",
"tokio",
"veil",
@ -3345,6 +3351,7 @@ dependencies = [
"google-cloud-auth",
"google-cloud-kms-v1",
"litellm-core-utils",
"litellm-http",
"litellm-python-compat",
"litellm-secrets-aws",
"litellm-secrets-azure",
@ -3392,6 +3399,7 @@ dependencies = [
"litellm-auth-azure",
"litellm-auth-types",
"litellm-core-utils",
"litellm-http",
"litellm-secrets-types",
"percent-encoding",
"reqwest 0.12.28",
@ -3411,6 +3419,7 @@ version = "0.1.0"
dependencies = [
"base64 0.22.1",
"litellm-core-utils",
"litellm-http",
"litellm-secrets-types",
"litellm-tracing",
"moka",
@ -3439,6 +3448,7 @@ dependencies = [
"litellm-auth-gcp",
"litellm-auth-types",
"litellm-core-utils",
"litellm-http",
"litellm-secrets-types",
"moka",
"percent-encoding",

View file

@ -7,4 +7,16 @@ disallowed-methods = [
{ path = "pyo3_async_runtimes::tokio::local_future_into_py", reason = "use litellm_host_python::run_async / run_async_value" },
{ path = "pyo3_async_runtimes::tokio::run", reason = "use litellm_host_python::run_sync / run_sync_value" },
{ path = "pyo3_async_runtimes::tokio::run_until_complete", reason = "use litellm_host_python::run_sync / run_sync_value" },
{ path = "reqwest::Client::new", reason = "take litellm_http::Client from HttpClientPool" },
{ path = "reqwest::Client::builder", reason = "HttpClientConfig owns client construction" },
{ path = "reqwest::ClientBuilder::danger_accept_invalid_certs", reason = "set HttpClientConfig::verify instead" },
{ path = "reqwest::ClientBuilder::identity", reason = "set HttpClientConfig::client_certificate instead" },
{ path = "reqwest::ClientBuilder::use_preconfigured_tls", reason = "HttpClientConfig owns the TLS configuration" },
]
# Every outbound client comes from litellm_http::HttpClientPool so it honors the host's TLS,
# proxy and timeout settings. Only crates/http builds one.
disallowed-types = [
{ path = "reqwest::Client", reason = "take litellm_http::Client from HttpClientPool; only crates/http builds one" },
{ path = "reqwest::ClientBuilder", reason = "HttpClientConfig owns client construction" },
]

View file

@ -22,5 +22,6 @@ aws-types = "1.4.0"
aws-smithy-runtime-api = "1.13.0"
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
reqwest.workspace = true
tokio.workspace = true

View file

@ -962,7 +962,7 @@ mod tests {
&no_env,
)
.await?;
let client = reqwest::Client::new();
let client = litellm_http::Client::plain_for_test();
let mut failures = Vec::new();
for region in ["us-west-2", "us-east-1"] {

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
litellm-http.workspace = true
litellm-auth-azure.workspace = true
litellm-auth-types.workspace = true
litellm-cache.workspace = true
@ -19,6 +20,7 @@ tokio.workspace = true
url.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
litellm-cache-response.workspace = true
litellm-cache-testing.workspace = true
rstest.workspace = true

View file

@ -31,7 +31,7 @@ impl<C: CacheCodec> AzureBlobCache<C> {
pub async fn connect(
account_url: &str,
container: &str,
http: reqwest::Client,
http: litellm_http::Client,
codec: C,
runtime: Handle,
) -> Result<Self, Error> {

View file

@ -8,7 +8,7 @@ use azure_core::{
use futures_util::TryStreamExt;
#[derive(Debug)]
pub struct ReqwestTransport(pub reqwest::Client);
pub struct ReqwestTransport(pub litellm_http::Client);
#[async_trait::async_trait]
impl HttpClient for ReqwestTransport {

View file

@ -31,7 +31,7 @@ async fn connect(server: &MockServer) -> AzureBlobCache<JsonCodec<Value>> {
None,
ClientOptions {
transport: Some(Transport::new(Arc::new(ReqwestTransport(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
)))),
..ClientOptions::default()
},

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
litellm-http.workspace = true
futures-util.workspace = true
litellm-auth-gcp.workspace = true
litellm-auth-types.workspace = true
@ -15,6 +16,7 @@ reqwest.workspace = true
tokio.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
litellm-cache-testing.workspace = true
rstest.workspace = true
serde_json.workspace = true

View file

@ -5,8 +5,8 @@ use litellm_cache::{
BaseCache, BatchCache, BatchEntry, CacheCodec, DisconnectCache, Error, ExactCacheContext,
FlushCache,
};
use litellm_http::Client;
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, percent_encode};
use reqwest::Client;
use crate::{GcpTokenSource, TokenSource};

View file

@ -96,7 +96,7 @@ async fn cache_exposes_its_configuration(#[future(awt)] server: MockServer) {
path_service_account: Some("/secrets/sa.json".into()),
..support::config(&server, Some("folder"))
},
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
litellm_cache::JsonCodec::<Value>::new(),
);
assert_eq!(cache.bucket_name(), "bucket");

View file

@ -29,7 +29,7 @@ pub fn cache_with_token(
) -> JsonGcsCache {
GcsCache::with_token_source(
config(server, gcs_path),
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
JsonCodec::new(),
token,
)

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
litellm-http.workspace = true
futures-util.workspace = true
litellm-cache.workspace = true
qdrant-client = { workspace = true, features = ["serde"] }
@ -17,6 +18,7 @@ tokio.workspace = true
uuid.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
futures-executor = "0.3"
litellm-cache-testing.workspace = true
rstest.workspace = true

View file

@ -1,7 +1,7 @@
use std::time::Duration;
use litellm_cache::{Error, semantic::Embedder};
use reqwest::Client;
use litellm_http::Client;
use serde_json::Value;
pub struct OpenAiEmbedder {

View file

@ -5,6 +5,10 @@ use std::{
use litellm_cache::{Error, semantic::Embedder};
use litellm_cache_qdrant_semantic::{OpenAiEmbedder, OpenAiEmbedderConfig};
use litellm_http::{
ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, Resolution,
media::PublicDnsResolver,
};
use rstest::rstest;
use serde_json::{Value, json};
use tokio::{
@ -104,7 +108,7 @@ fn config(base: String, timeout: Option<Duration>) -> OpenAiEmbedderConfig {
async fn posts_embeddings_request_and_parses_vector() {
let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await;
let embedder = OpenAiEmbedder::new(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
config(
format!("{}/", server.base_url()),
Some(Duration::from_secs(1)),
@ -156,14 +160,17 @@ async fn status_timeout_and_body_errors_are_unavailable(
) {
let server =
TestHttpServer::response_after(status, body, Duration::from_millis(delay_ms)).await;
let embedder = OpenAiEmbedder::new(reqwest::Client::new(), config(server.base_url(), timeout));
let embedder = OpenAiEmbedder::new(
litellm_http::Client::plain_for_test(),
config(server.base_url(), timeout),
);
assert_eq!(embedder.async_embed("hello", None).await, expected);
}
#[rstest]
fn sync_embedding_is_unsupported() {
let embedder = OpenAiEmbedder::new(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
config("http://127.0.0.1:9".to_owned(), None),
);
assert_eq!(
@ -176,9 +183,12 @@ fn sync_embedding_is_unsupported() {
#[tokio::test]
async fn uses_the_injected_client() {
let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await;
let client = reqwest::Client::builder()
.user_agent("litellm-embedder-test")
.build()
let config_with_agent = HttpClientConfig {
user_agent: Some("litellm-embedder-test".into()),
..Resolution::from(&HttpSettings::default()).config
};
let client = HttpClientPool::new(Arc::new(PublicDnsResolver))
.client(&config_with_agent, ClientVariant::Provider)
.unwrap();
let embedder = OpenAiEmbedder::new(client, config(server.base_url(), None));
assert_eq!(

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
litellm-http.workspace = true
litellm-cache.workspace = true
litellm-auth-aws.workspace = true
aws-sdk-s3 = { version = "1.146.1", default-features = false, features = ["rustls", "rt-tokio"] }
@ -19,6 +20,7 @@ reqwest.workspace = true
tokio.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
litellm-cache-testing.workspace = true
rstest.workspace = true
wiremock = "0.6.5"

View file

@ -42,7 +42,12 @@ pub struct S3Cache<C: CacheCodec> {
}
impl<C: CacheCodec> S3Cache<C> {
pub fn new(config: S3CacheConfig, http: reqwest::Client, codec: C, runtime: Handle) -> Self {
pub fn new(
config: S3CacheConfig,
http: litellm_http::Client,
codec: C,
runtime: Handle,
) -> Self {
let endpoint_url: Option<String> = config.endpoint.map(|endpoint| endpoint.url);
let base = aws_sdk_s3::Config::builder()
.behavior_version(BehaviorVersion::latest())

View file

@ -9,7 +9,7 @@ use aws_smithy_runtime_api::client::{
use aws_smithy_types::body::SdkBody;
#[derive(Clone, Debug)]
pub(crate) struct ReqwestHttpClient(pub(crate) reqwest::Client);
pub(crate) struct ReqwestHttpClient(pub(crate) litellm_http::Client);
impl HttpClient for ReqwestHttpClient {
fn http_connector(

View file

@ -32,7 +32,12 @@ pub fn config(endpoint: &str) -> S3CacheConfig {
}
pub fn cache_with(config: S3CacheConfig, runtime: Handle) -> JsonS3Cache {
S3Cache::new(config, reqwest::Client::new(), JsonCodec::new(), runtime)
S3Cache::new(
config,
litellm_http::Client::plain_for_test(),
JsonCodec::new(),
runtime,
)
}
pub fn cache(endpoint: &str) -> JsonS3Cache {

View file

@ -36,6 +36,7 @@ url.workspace = true
veil.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
litellm-auth-gcp.workspace = true
litellm-llms = { workspace = true, features = ["test-support"] }
rstest.workspace = true

View file

@ -1,13 +0,0 @@
use std::{sync::OnceLock, time::Duration};
use crate::constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS;
pub(super) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS))
.build()
.unwrap_or_else(|_| reqwest::Client::new())
})
}

View file

@ -1,10 +1,16 @@
use litellm_http::request::truncate_error_body;
use std::time::Duration;
use litellm_http::{Client, request::truncate_error_body};
use serde_json::Value;
use super::{Error, client::http_client};
use crate::audio_transcription::types::ProviderAudioTranscriptionRequest;
use super::Error;
use crate::{
audio_transcription::types::ProviderAudioTranscriptionRequest,
constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS,
};
pub async fn execute_audio_transcription_provider_call(
http: &Client,
request: ProviderAudioTranscriptionRequest,
) -> Result<Value, Error> {
let response = crate::outbound::outbound_request::<Error>(
@ -12,11 +18,15 @@ pub async fn execute_audio_transcription_provider_call(
request.url.clone(),
request.upstream_headers.clone(),
&request.body,
request.timeout,
Some(
request
.timeout
.unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)),
),
&request.optional_params,
)
.await?
.send(http_client())
.send(http)
.await
.map_err(|error| {
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))

View file

@ -1,16 +1,21 @@
mod error;
pub mod types;
pub use error::Error;
mod client;
mod handler;
mod prepare;
pub use handler::execute_audio_transcription_provider_call;
use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool};
pub use prepare::prepare_audio_transcription_provider_call;
use serde_json::Value;
use crate::audio_transcription::types::AudioTranscriptionRequest;
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
.await
pub async fn audio_transcription(
pool: &HttpClientPool,
config: &HttpClientConfig,
request: AudioTranscriptionRequest<'_>,
) -> Result<Value, Error> {
let request = prepare_audio_transcription_provider_call(request)?;
let http = pool.client(config, ClientVariant::Provider)?;
execute_audio_transcription_provider_call(&http, request).await
}

View file

@ -1,14 +0,0 @@
use std::{sync::OnceLock, time::Duration};
use crate::constants::{CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS, CHAT_COMPLETIONS_TIMEOUT_SECS};
pub(super) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS))
.connect_timeout(Duration::from_secs(CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS))
.build()
.unwrap_or_else(|_| reqwest::Client::new())
})
}

View file

@ -1,20 +1,24 @@
use litellm_http::{outbound::OutboundRequest, request::truncate_error_body};
use std::time::Duration;
use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body};
use litellm_llms::base_llm::chat::transformation::ProviderChatResponseData;
use litellm_types::utils::ChatCompletionsResponse;
use serde_json::Value;
use super::{Error, client::http_client, prepare::prepare_provider_request};
use crate::chat_completions::types::{
ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
use super::{Error, prepare::prepare_provider_request};
use crate::{
chat_completions::types::{ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest},
constants::CHAT_COMPLETIONS_TIMEOUT_SECS,
};
pub(super) async fn execute_chat_completions_provider_call(
http: &Client,
request: ResolvedChatCompletionsRequest<'_>,
) -> Result<ChatCompletionsResponse, Error> {
let request = prepare_provider_request(request)?;
let outbound = outbound_request(&request).await?;
let response = outbound.send(http_client()).await.map_err(|err| {
let response = outbound.send(http).await.map_err(|err| {
// Failing to establish the connection means the request never went out,
// so the host can still serve it. Everything else here, a timeout
// above all, may have reached the provider and been answered.
@ -72,7 +76,11 @@ pub(super) async fn outbound_request(
request.url.clone(),
request.upstream_headers.clone(),
&request.body,
request.timeout,
Some(
request
.timeout
.unwrap_or(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS)),
),
&request.optional_params,
)
.await

View file

@ -9,11 +9,11 @@
mod error;
pub mod types;
pub use error::Error;
mod client;
mod common_utils;
pub(crate) mod handler;
mod prepare;
use handler::execute_chat_completions_provider_call;
use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool};
use litellm_types::utils::ChatCompletionsResponse;
use prepare::{parse_messages, resolve_provider_config, resolve_request};
use serde_json::{Map, Value};
@ -21,9 +21,13 @@ use serde_json::{Map, Value};
use crate::chat_completions::types::ChatCompletionsRequest;
pub async fn chat_completions(
pool: &HttpClientPool,
config: &HttpClientConfig,
request: ChatCompletionsRequest<'_>,
) -> Result<ChatCompletionsResponse, Error> {
execute_chat_completions_provider_call(resolve_request(request)?).await
let request = resolve_request(request)?;
let http = pool.client(config, ClientVariant::Provider)?;
execute_chat_completions_provider_call(&http, request).await
}
/// Whether the core would accept this request, without resolving credentials or

View file

@ -5,9 +5,6 @@ pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com";
/// timeout from the caller still overrides this on the request builder.
pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
/// Connect timeout for Anthropic Messages provider calls, in seconds.
pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10;
/// Provider name used for Anthropic Messages when a deployment's provider model
/// does not carry an explicit provider prefix.
pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
@ -16,9 +13,6 @@ pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
/// seconds. Mirrors the Python chat completions default.
pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600;
/// Connect timeout for chat completions provider calls, in seconds.
pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600;
/// `object` field every non-streaming chat completion response carries.

View file

@ -1,14 +0,0 @@
use std::{sync::OnceLock, time::Duration};
use crate::constants::{MESSAGES_CONNECT_TIMEOUT_SECS, MESSAGES_TIMEOUT_SECS};
pub(super) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(MESSAGES_TIMEOUT_SECS))
.connect_timeout(Duration::from_secs(MESSAGES_CONNECT_TIMEOUT_SECS))
.build()
.unwrap_or_else(|_| reqwest::Client::new())
})
}

View file

@ -17,6 +17,8 @@ pub enum Error {
#[error(transparent)]
Auth(#[from] litellm_auth::Error),
#[error(transparent)]
Client(#[from] litellm_http::Error),
#[error(transparent)]
Transport(#[from] litellm_http::transport::Error),
#[error(transparent)]
Headers(#[from] litellm_http::request::HeaderError),

View file

@ -5,13 +5,15 @@ use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMes
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use serde_json::Value;
use super::{Error, client::http_client, common_utils::truncate_error_body};
use super::{Error, common_utils::truncate_error_body};
use crate::constants::MESSAGES_TIMEOUT_SECS;
pub(super) fn network(error: reqwest::Error) -> Error {
Error::Transport(TransportError::Network(error.to_string()))
}
pub(super) async fn send(
http: &litellm_http::Client,
url: &str,
headers: &[(String, String)],
body: &Value,
@ -20,13 +22,11 @@ pub(super) async fn send(
let encoded = serde_json::to_vec(body)
.map_err(|err| Error::InvalidRequest(format!("failed to encode messages body: {err}")))?;
let builder = headers.iter().fold(
http_client().post(url).body(encoded),
http.post(url)
.body(encoded)
.timeout(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))),
|builder, (key, value)| builder.header(key, value),
);
let builder = match timeout {
Some(duration) => builder.timeout(duration),
None => builder,
};
http_request(builder).await.map_err(network)
}

View file

@ -7,13 +7,13 @@
mod error;
pub mod types;
pub use error::Error;
mod client;
mod common_utils;
mod handler;
mod prepare;
pub mod route;
use std::sync::Arc;
use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool};
use litellm_secrets::source::EnvironmentSecrets;
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine};
@ -21,7 +21,11 @@ use serde_json::Value;
use crate::messages::types::MessagesRequest;
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
pub async fn messages(
pool: &HttpClientPool,
config: &HttpClientConfig,
request: MessagesRequest<'_>,
) -> Result<AnthropicMessagesResponse, Error> {
let Value::Object(body) = request.body else {
return Err(Error::InvalidRequest(
"messages body must be an object".into(),
@ -38,8 +42,15 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesR
timeout: request.timeout,
shaping: request.shaping,
};
let secrets = Arc::new(EnvironmentSecrets::python_compatible());
match litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await? {
let secrets = Arc::new(EnvironmentSecrets::python_compatible(
pool.client(config, ClientVariant::Provider)?,
));
match litellm_host::run::run(
messages_machine(pool, config, secrets)?,
&LocalMessagesHost::new(call),
)
.await?
{
MessagesOutput::Message(message) => Ok(*message),
MessagesOutput::Streamed => Err(Error::Unsupported(
"streamed responses need a streaming host",

View file

@ -12,6 +12,7 @@ use litellm_host::{
machine::{CallMachine, HostChannel, MachineFault},
protocol::Protocol,
};
use litellm_http::{Client, ClientVariant, HttpClientConfig, HttpClientPool};
use litellm_secrets::source::SecretSource;
use litellm_types::{
llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
@ -108,12 +109,20 @@ impl Host<Messages> for LocalMessagesHost {
}
}
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
CallMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
pub fn messages_machine(
pool: &HttpClientPool,
config: &HttpClientConfig,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesMachine, litellm_http::Error> {
let http = pool.client(config, ClientVariant::Provider)?;
Ok(CallMachine::new(move |host| {
Box::pin(execute(host, http.clone(), secrets.clone()))
}))
}
async fn execute(
host: MessagesHost,
http: Client,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesOutput, Error> {
let call = host.project().await?;
@ -164,7 +173,7 @@ async fn execute(
context,
)
.await?;
let response = send(&wire.url, &wire.headers, &wire.body, request.timeout).await?;
let response = send(&http, &wire.url, &wire.headers, &wire.body, request.timeout).await?;
if !response.status().is_success() {
return Err(provider_error(response).await);
}

View file

@ -110,7 +110,10 @@ mod tests {
}
fn client() -> OcrClient {
OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new())
OcrClient::for_test(
litellm_http::Client::plain_for_test(),
litellm_http::Client::no_redirect_for_test(),
)
}
fn request(model: &str, base: &str, document: Value, options: Value) -> LiteLLMOcrRequest {

View file

@ -10,6 +10,10 @@ use support::*;
const MODEL: &str = "mistral.voxtral-mini-3b-2507";
async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
audio_transcription(&http_pool(), &http_config(), request).await
}
fn transcript_response(text: &str) -> ResponseTemplate {
json_response(json!({"output": {"message": {"content": [{"text": text}]}}}))
}
@ -47,7 +51,7 @@ async fn bedrock_converse_request_is_signed_for_the_requested_region(
let upstream = upstream([transcript_response("hello")]).await;
let base = upstream.uri();
let response = audio_transcription(AudioTranscriptionRequest {
let response = transcribe(AudioTranscriptionRequest {
api_base: Some(&base),
optional_params: aws_params(region),
..request
@ -79,7 +83,7 @@ async fn the_provider_can_come_from_the_model_prefix(request: AudioTranscription
let base = upstream.uri();
let model = format!("bedrock/{MODEL}");
audio_transcription(AudioTranscriptionRequest {
transcribe(AudioTranscriptionRequest {
model: &model,
custom_llm_provider: None,
api_base: Some(&base),
@ -110,7 +114,7 @@ async fn audio_and_transcription_params_reach_the_converse_body(
])
.collect();
audio_transcription(AudioTranscriptionRequest {
transcribe(AudioTranscriptionRequest {
audio: json!({"data": "AQI=", "format": format}),
api_base: Some(&base),
optional_params,
@ -142,7 +146,7 @@ async fn invalid_audio_is_rejected_before_sending(
let upstream = upstream([transcript_response("hello")]).await;
let base = upstream.uri();
let error = audio_transcription(AudioTranscriptionRequest {
let error = transcribe(AudioTranscriptionRequest {
audio,
api_base: Some(&base),
..request
@ -174,7 +178,7 @@ async fn unsupported_providers_are_rejected_before_sending(
#[case] provider: Option<&'static str>,
#[case] reported: &str,
) {
let error = audio_transcription(AudioTranscriptionRequest {
let error = transcribe(AudioTranscriptionRequest {
model,
custom_llm_provider: provider,
api_base: Some(UNREACHABLE_BASE),
@ -189,7 +193,7 @@ async fn unsupported_providers_are_rejected_before_sending(
#[rstest]
#[tokio::test]
async fn a_non_string_extra_header_is_rejected(request: AudioTranscriptionRequest<'static>) {
let error = audio_transcription(AudioTranscriptionRequest {
let error = transcribe(AudioTranscriptionRequest {
extra_headers: Some(Map::from_iter([("x-count".to_string(), json!(3))])),
api_base: Some(UNREACHABLE_BASE),
..request
@ -212,7 +216,7 @@ async fn an_upstream_error_keeps_its_status_and_body(
upstream([ResponseTemplate::new(status).set_body_string("upstream said no")]).await;
let base = upstream.uri();
let error = audio_transcription(AudioTranscriptionRequest {
let error = transcribe(AudioTranscriptionRequest {
api_base: Some(&base),
..request
})
@ -239,7 +243,7 @@ async fn an_unreadable_success_body_is_an_invalid_response(
let upstream = upstream([response]).await;
let base = upstream.uri();
let error = audio_transcription(AudioTranscriptionRequest {
let error = transcribe(AudioTranscriptionRequest {
api_base: Some(&base),
..request
})

View file

@ -4,6 +4,7 @@ use litellm_core::chat_completions::{
Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest,
};
use litellm_http::transport::Error as TransportError;
use litellm_types::utils::ChatCompletionsResponse;
use rstest::{fixture, rstest};
use serde_json::{Map, Value, json};
use wiremock::ResponseTemplate;
@ -13,6 +14,10 @@ use support::*;
const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
async fn complete(request: ChatCompletionsRequest<'_>) -> Result<ChatCompletionsResponse, Error> {
chat_completions(&http_pool(), &http_config(), request).await
}
fn object(value: Value) -> Map<String, Value> {
let Value::Object(map) = value else {
panic!("expected a json object, got {value}");
@ -50,7 +55,7 @@ async fn anthropic_round_trip_translates_the_conversation_and_normalizes_the_res
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
let response = chat_completions(ChatCompletionsRequest {
let response = complete(ChatCompletionsRequest {
messages: json!([
{"role": "system", "content": "be terse"},
{"role": "user", "content": "hi"}
@ -90,7 +95,7 @@ async fn the_deployment_key_replaces_a_caller_supplied_x_api_key(
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
chat_completions(ChatCompletionsRequest {
complete(ChatCompletionsRequest {
api_base: Some(&base),
extra_headers: Some(object(
json!({"x-api-key": "caller-key", "x-trace": "kept"}),
@ -116,7 +121,7 @@ async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsReq
.await;
let base = upstream.uri();
let response = chat_completions(ChatCompletionsRequest {
let response = complete(ChatCompletionsRequest {
model: "bedrock/anthropic.claude-sonnet-4-5",
optional_params: object(json!({
"aws_access_key_id": "access-key",
@ -167,7 +172,7 @@ async fn a_response_it_cannot_normalize_is_reported_as_already_sent(
let upstream = upstream([anthropic_response(body)]).await;
let base = upstream.uri();
let error = chat_completions(ChatCompletionsRequest {
let error = complete(ChatCompletionsRequest {
api_base: Some(&base),
..request
})
@ -188,7 +193,7 @@ async fn an_upstream_error_status_keeps_its_code_and_body(
let upstream = upstream([ResponseTemplate::new(status).set_body_string("slow down")]).await;
let base = upstream.uri();
let error = chat_completions(ChatCompletionsRequest {
let error = complete(ChatCompletionsRequest {
api_base: Some(&base),
..request
})
@ -210,7 +215,7 @@ async fn an_upstream_error_status_keeps_its_code_and_body(
async fn a_connection_that_is_never_established_declines_instead_of_failing(
request: ChatCompletionsRequest<'static>,
) {
let error = chat_completions(ChatCompletionsRequest {
let error = complete(ChatCompletionsRequest {
api_base: Some(UNREACHABLE_BASE),
..request
})
@ -232,7 +237,7 @@ async fn a_timeout_after_sending_is_not_a_pre_send_decline(
upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await;
let base = upstream.uri();
let error = chat_completions(ChatCompletionsRequest {
let error = complete(ChatCompletionsRequest {
api_base: Some(&base),
timeout: Some(Duration::from_millis(100)),
..request
@ -307,7 +312,7 @@ async fn a_declined_request_fails_the_call_before_sending(
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
let error = chat_completions(ChatCompletionsRequest {
let error = complete(ChatCompletionsRequest {
optional_params: object(json!({"stream": true})),
api_base: Some(&base),
..request

View file

@ -78,7 +78,7 @@ impl Host<Messages> for RecordingHost {
}
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
}
fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall {

View file

@ -2,9 +2,11 @@ use std::{sync::Arc, time::Duration};
use litellm_core::messages::{
Error,
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
route::{LocalMessagesHost, MessagesCall, MessagesMachine, MessagesOutput, messages_machine},
types::MessagesShaping,
};
use litellm_http::{HttpSettings, Resolution};
use litellm_secrets::source::SecretSource;
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use rstest::fixture;
use serde_json::{Map, Value, json};
@ -75,11 +77,16 @@ fn headers<'a>(pairs: impl IntoIterator<Item = (&'a str, &'a str)>) -> Option<Ma
)
}
fn machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
messages_machine(&http_pool(), &http_config(), secrets)
.expect("default HTTP settings build a client")
}
async fn run_with(
secrets: Arc<RecordingSecrets>,
call: MessagesCall,
) -> Result<MessagesOutput, Error> {
litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await
litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await
}
/// Runs the route with a secret source that knows nothing, so no environment leaks in.

View file

@ -190,29 +190,40 @@ fn facade_request(body: Value, api_base: &str) -> MessagesRequest<'_> {
}
#[tokio::test]
async fn the_facade_runs_the_route_in_process() {
async fn the_facade_sends_through_the_injected_http_pool_configuration() {
let upstream = upstream([message_response()]).await;
let base = upstream.uri();
let settings = HttpSettings {
user_agent: Some("host-owned/1".into()),
..HttpSettings::default()
};
let message = messages(facade_request(
json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}),
&base,
))
let message = messages(
&http_pool(),
&Resolution::from(&settings).config,
facade_request(
json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}),
&base,
),
)
.await
.expect("messages request succeeds");
assert_eq!(message.id, "msg_1");
assert_eq!(
only_request(&upstream).await.header("x-api-key"),
Some("sk-ant")
);
let sent = only_request(&upstream).await;
assert_eq!(sent.header("x-api-key"), Some("sk-ant"));
assert_eq!(sent.header("user-agent"), Some("host-owned/1"));
}
#[tokio::test]
async fn the_facade_rejects_a_body_that_is_not_an_object() {
let error = messages(facade_request(json!([]), UNREACHABLE_BASE))
.await
.expect_err("a non-object body is rejected");
let error = messages(
&http_pool(),
&http_config(),
facade_request(json!([]), UNREACHABLE_BASE),
)
.await
.expect_err("a non-object body is rejected");
assert_eq!(
error,

View file

@ -87,7 +87,7 @@ fn sse_response() -> ResponseTemplate {
}
async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Error> {
litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
}
#[rstest]

View file

@ -4,6 +4,7 @@ use litellm_core::ocr::{
types::LiteLLMOcrRequest,
wire::{OcrWireRequest, decode_request},
};
use litellm_http::Client;
use litellm_llms::base_llm::ocr::{
error::Error,
handler::OcrClient,
@ -37,11 +38,7 @@ fn object(value: Value) -> Map<String, Value> {
}
fn ocr_client() -> OcrClient {
let document_http = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test document client builds");
OcrClient::for_test(reqwest::Client::new(), document_http)
OcrClient::for_test(Client::plain_for_test(), Client::no_redirect_for_test())
}
async fn perform(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {

View file

@ -1,10 +1,7 @@
use std::sync::Arc;
use litellm_auth_gcp::VertexAuth;
use litellm_http::{
HttpClientPool, HttpSettings, Resolution,
media::{PublicDnsResolver, UrlPolicy},
};
use litellm_http::{HttpSettings, Resolution, media::UrlPolicy};
use litellm_llms::{
base_llm::ocr::{
settings::OcrSettings,
@ -192,12 +189,16 @@ async fn the_client_uses_the_injected_http_pool_configuration() {
..HttpSettings::default()
};
let client = OcrClient::new(
&HttpClientPool::new(Arc::new(PublicDnsResolver)),
&http_pool(),
&Resolution::from(&settings).config,
UrlPolicy::default(),
VertexAuth::default(),
OcrSettings::default(),
Arc::new(litellm_secrets::source::EnvironmentSecrets::default()),
Arc::new(
litellm_secrets::source::EnvironmentSecrets::python_compatible(
litellm_http::Client::plain_for_test(),
),
),
)
.unwrap();

View file

@ -3,9 +3,12 @@
#![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset
use std::sync::Mutex;
use std::sync::{Arc, Mutex};
use futures_util::future::BoxFuture;
use litellm_http::{
HttpClientConfig, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver,
};
use litellm_secrets::{SecretValue, source::SecretSource};
use serde_json::Value;
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
@ -13,6 +16,14 @@ use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
/// A port nothing listens on, for calls that must fail before any request is sent.
pub const UNREACHABLE_BASE: &str = "http://127.0.0.1:1";
pub fn http_pool() -> HttpClientPool {
HttpClientPool::new(Arc::new(PublicDnsResolver))
}
pub fn http_config() -> HttpClientConfig {
Resolution::from(&HttpSettings::default()).config
}
/// Starts an upstream that answers its n-th request with the n-th response and 404s after.
pub async fn upstream(responses: impl IntoIterator<Item = ResponseTemplate>) -> MockServer {
let server = MockServer::start().await;

View file

@ -22,5 +22,7 @@ veil.workspace = true
webpki-roots.workspace = true
[dev-dependencies]
rcgen = "0.14.10"
tempfile.workspace = true
rstest.workspace = true
tokio.workspace = true

View file

@ -0,0 +1,38 @@
use std::ops::Deref;
#[derive(Clone, Debug)]
pub struct Client(reqwest::Client);
impl Client {
pub(crate) fn new(client: reqwest::Client) -> Self {
Self(client)
}
#[cfg(any(test, feature = "test-support"))]
pub fn for_test(client: reqwest::Client) -> Self {
Self(client)
}
#[cfg(any(test, feature = "test-support"))]
pub fn plain_for_test() -> Self {
Self(reqwest::Client::new())
}
#[cfg(any(test, feature = "test-support"))]
pub fn no_redirect_for_test() -> Self {
Self(
reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("a client without TLS or proxy settings builds"),
)
}
}
impl Deref for Client {
type Target = reqwest::Client;
fn deref(&self) -> &reqwest::Client {
&self.0
}
}

View file

@ -18,10 +18,16 @@ pub enum Verify {
BuiltInRoots,
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub enum ClientIdentity {
Pem(PathBuf),
Split { certificate: PathBuf, key: PathBuf },
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub struct HttpClientConfig {
pub verify: Verify,
pub client_certificate: Option<PathBuf>,
pub client_certificate: Option<ClientIdentity>,
pub key_exchange_group: Option<KeyExchangeGroup>,
pub tls12_cipher_suites: Option<Vec<Tls12CipherSuite>>,
pub force_ipv4: bool,
@ -67,7 +73,7 @@ impl From<&HttpSettings> for Resolution {
Self {
config: HttpClientConfig {
verify: Verify::from(settings),
client_certificate: settings.ssl_certificate.clone(),
client_certificate: settings.ssl_certificate.clone().map(ClientIdentity::Pem),
key_exchange_group: curve.clone().ok().flatten(),
tls12_cipher_suites: ciphers.tls12_cipher_suites,
force_ipv4: settings.force_ipv4,
@ -276,7 +282,7 @@ mod tests {
config,
HttpClientConfig {
verify: Verify::BuiltInRoots,
client_certificate: Some("/client.pem".into()),
client_certificate: Some(ClientIdentity::Pem("/client.pem".into())),
key_exchange_group: None,
tls12_cipher_suites: None,
force_ipv4: true,

View file

@ -1,3 +1,10 @@
#![allow(
clippy::disallowed_types,
clippy::disallowed_methods,
reason = "this crate is the one place reqwest clients are built"
)]
mod client;
mod config;
mod error;
pub mod media;
@ -9,7 +16,8 @@ mod settings;
mod tls;
pub mod transport;
pub use config::{HttpClientConfig, Resolution, Verify};
pub use client::Client;
pub use config::{ClientIdentity, HttpClientConfig, Resolution, Verify};
pub use error::{Error, TlsSource};
pub use pool::{ClientVariant, HttpClientPool};
pub use proxy::EnvironmentProxies;

View file

@ -12,7 +12,7 @@ use reqwest::{
dns::{Addrs, Name, Resolve, Resolving},
};
use crate::{ClientVariant, HttpClientConfig, HttpClientPool};
use crate::{Client, ClientVariant, HttpClientConfig, HttpClientPool};
#[derive(Debug, thiserror::Error)]
pub enum Error {
@ -93,8 +93,8 @@ type ProxyMatch = Arc<dyn Fn(&Url) -> bool + Send + Sync>;
#[derive(Clone)]
pub struct MediaFetcher {
pinned: reqwest::Client,
unpinned: reqwest::Client,
pinned: Client,
unpinned: Client,
uses_proxy: ProxyMatch,
address_resolver: Arc<dyn AddressResolver>,
url_policy: UrlPolicy,
@ -154,7 +154,7 @@ impl MediaFetcher {
}
#[cfg(any(test, feature = "test-support"))]
pub fn for_test(client: reqwest::Client) -> Self {
pub fn for_test(client: Client) -> Self {
Self {
pinned: client.clone(),
unpinned: client,
@ -230,7 +230,7 @@ impl MediaFetcher {
}
}
async fn client_for(&self, url: &Url) -> Result<&reqwest::Client, Error> {
async fn client_for(&self, url: &Url) -> Result<&Client, Error> {
if !self.url_policy.validate {
return Ok(&self.unpinned);
}
@ -520,10 +520,7 @@ mod tests {
b"HTTP/1.1 200 OK\r\nContent-Type: application/pdf; charset=binary\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc",
)
.await;
let client = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test client builds");
let client = Client::no_redirect_for_test();
let media = MediaFetcher::for_test(client)
.fetch(url, policy(3, 0))
.await
@ -539,10 +536,7 @@ mod tests {
b"HTTP/1.1 200 OK\r\nContent-Type: application/pdf\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc",
)
.await;
let client = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test client builds");
let client = Client::no_redirect_for_test();
let error = MediaFetcher::for_test(client)
.fetch(url, policy(2, 0))
.await
@ -557,10 +551,7 @@ mod tests {
b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n2\r\nab\r\n2\r\ncd\r\n0\r\n\r\n",
)
.await;
let client = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test client builds");
let client = Client::no_redirect_for_test();
let error = MediaFetcher::for_test(client)
.fetch(url, policy(3, 0))
.await

View file

@ -107,7 +107,7 @@ impl OutboundRequest {
self.timeout
}
pub async fn send(self, client: &reqwest::Client) -> Result<reqwest::Response, reqwest::Error> {
pub async fn send(self, client: &crate::Client) -> Result<reqwest::Response, reqwest::Error> {
let builder = with_headers(
client.post(&self.url).body(self.body),
&self.headers,

View file

@ -6,7 +6,7 @@ use std::{
use reqwest::dns::Resolve;
use crate::{config::HttpClientConfig, error::Error, proxy::EnvironmentProxies};
use crate::{client::Client, config::HttpClientConfig, error::Error, proxy::EnvironmentProxies};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum ClientVariant {
@ -48,7 +48,7 @@ impl HttpClientPool {
&self,
config: &HttpClientConfig,
variant: ClientVariant,
) -> Result<reqwest::Client, Error> {
) -> Result<Client, Error> {
let effective = match variant {
ClientVariant::Media => HttpClientConfig {
client_certificate: None,
@ -65,7 +65,7 @@ impl HttpClientPool {
if let Some(pooled) = self.lock().get(&key)
&& pooled.built_at.elapsed() < self.ttl
{
return Ok(pooled.client.clone());
return Ok(Client::new(pooled.client.clone()));
}
let client = self
.apply(variant, reqwest::ClientBuilder::try_from(&key.0)?)
@ -77,7 +77,7 @@ impl HttpClientPool {
built_at: Instant::now(),
},
);
Ok(client)
Ok(Client::new(client))
}
fn lock(&self) -> MutexGuard<'_, Clients> {
@ -116,7 +116,7 @@ mod tests {
};
use super::*;
use crate::{HttpSettings, Resolution, Verify};
use crate::{ClientIdentity, HttpSettings, Resolution, Verify};
struct FixedResolver(SocketAddr);
@ -288,7 +288,9 @@ mod tests {
fn media_variant_never_loads_the_client_certificate() {
let pool = pool();
let with_identity = HttpClientConfig {
client_certificate: Some(std::env::temp_dir().join("litellm-http-absent-client.pem")),
client_certificate: Some(ClientIdentity::Pem(
std::env::temp_dir().join("litellm-http-absent-client.pem"),
)),
..config("a")
};
assert!(

View file

@ -8,7 +8,7 @@ use rustls::{
};
use crate::{
config::{HttpClientConfig, Verify},
config::{ClientIdentity, HttpClientConfig, Verify},
error::{Error, TlsSource},
};
@ -203,11 +203,15 @@ impl TryFrom<&HttpClientConfig> for ClientConfig {
};
let mut tls = match &config.client_certificate {
None => verified.with_no_client_auth(),
Some(path) => {
let (chain, key) = identity(path, TlsSource::ClientIdentity)?;
Some(identity) => {
let (certificate, key) = match identity {
ClientIdentity::Pem(path) => (path, path),
ClientIdentity::Split { certificate, key } => (certificate, key),
};
let (chain, private_key) = client_identity(certificate, key)?;
verified
.with_client_auth_cert(chain, key)
.map_err(|error| invalid_pem(path, TlsSource::ClientIdentity, error))?
.with_client_auth_cert(chain, private_key)
.map_err(|error| invalid_pem(key, TlsSource::ClientIdentity, error))?
}
};
tls.alpn_protocols = if config.http2 {
@ -233,17 +237,18 @@ fn bundle_roots(path: &Path, source: TlsSource) -> Result<RootCertStore, Error>
Ok(store)
}
fn identity(
path: &Path,
source: TlsSource,
fn client_identity(
certificate: &Path,
key: &Path,
) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>), Error> {
let chain = certificates(path, source)?;
let source = TlsSource::ClientIdentity;
let chain = certificates(certificate, source)?;
if chain.is_empty() {
return Err(invalid_pem(path, source, "no certificates found"));
return Err(invalid_pem(certificate, source, "no certificates found"));
}
let key = PrivateKeyDer::from_pem_slice(&read(path, source)?)
.map_err(|error| invalid_pem(path, source, error))?;
Ok((chain, key))
let private_key = PrivateKeyDer::from_pem_slice(&read(key, source)?)
.map_err(|error| invalid_pem(key, source, error))?;
Ok((chain, private_key))
}
fn certificates(path: &Path, source: TlsSource) -> Result<Vec<CertificateDer<'static>>, Error> {
@ -405,7 +410,7 @@ mod tests {
)
.unwrap();
let result = ClientConfig::try_from(&HttpClientConfig {
client_certificate: Some(path.clone()),
client_certificate: Some(ClientIdentity::Pem(path.clone())),
..config(HttpSettings::default())
})
.map(drop);
@ -419,4 +424,29 @@ mod tests {
}) if reported == path
));
}
#[test]
fn split_client_identity_reads_the_key_from_its_own_file() {
let identity = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap();
let directory = tempfile::tempdir().unwrap();
let certificate = directory.path().join("client.crt");
let key = directory.path().join("client.key");
std::fs::write(&certificate, identity.cert.pem()).unwrap();
std::fs::write(&key, identity.signing_key.serialize_pem()).unwrap();
let split = ClientConfig::try_from(&HttpClientConfig {
client_certificate: Some(ClientIdentity::Split {
certificate: certificate.clone(),
key,
}),
..config(HttpSettings::default())
});
let combined = ClientConfig::try_from(&HttpClientConfig {
client_certificate: Some(ClientIdentity::Pem(certificate)),
..config(HttpSettings::default())
});
assert!(split.unwrap().client_auth_cert_resolver.has_certs());
assert!(combined.is_err());
}
}

View file

@ -35,6 +35,7 @@ tokio = { workspace = true, features = ["sync"] }
url.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
aws-smithy-eventstream = "=0.61.4"
aws-smithy-types = "1.6.1"
rstest.workspace = true

View file

@ -450,7 +450,7 @@ fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result<i64, E
}
async fn read_operation_response(
http_client: &reqwest::Client,
http_client: &litellm_http::Client,
response: reqwest::Response,
original_url: &str,
headers: &[(String, String)],
@ -489,7 +489,7 @@ async fn read_operation_response(
}
async fn poll_operation(
http_client: &reqwest::Client,
http_client: &litellm_http::Client,
url: Url,
headers: &[(String, String)],
connection: &OcrConnection,

View file

@ -184,16 +184,21 @@ mod tests {
reqwest::header::AUTHORIZATION,
reqwest::header::HeaderValue::from_static("Bearer provider-secret"),
);
let provider_http = reqwest::Client::builder()
.default_headers(provider_headers)
.build()
.unwrap();
let document_http = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.unwrap();
let client =
crate::base_llm::ocr::handler::OcrClient::for_test(provider_http, document_http);
#[expect(
clippy::disallowed_methods,
clippy::disallowed_types,
reason = "the pool has no default-header setting to stand in for provider credentials"
)]
let provider_http = litellm_http::Client::for_test(
reqwest::Client::builder()
.default_headers(provider_headers)
.build()
.unwrap(),
);
let client = crate::base_llm::ocr::handler::OcrClient::for_test(
provider_http,
litellm_http::Client::no_redirect_for_test(),
);
let converted = inline_remote_document(
client.document_fetcher(),
OcrDocument::ImageUrl {

View file

@ -5,7 +5,7 @@ use futures_util::future::BoxFuture;
use litellm_auth_gcp::VertexAuth;
use litellm_host::event::WireRequest;
use litellm_http::{
ClientVariant, HttpClientConfig, HttpClientPool,
Client, ClientVariant, HttpClientConfig, HttpClientPool,
media::{MediaFetcher, UrlPolicy},
outbound::{OutboundRequest, RequestSigner},
transport,
@ -33,8 +33,8 @@ pub trait CallHooks<E>: Send + Sync {
#[derive(Clone)]
pub struct OcrClient {
provider_http: reqwest::Client,
polling_http: reqwest::Client,
provider_http: Client,
polling_http: Client,
document_fetcher: MediaFetcher,
vertex_auth: VertexAuth,
settings: OcrSettings,
@ -60,11 +60,11 @@ impl OcrClient {
})
}
pub fn provider_http(&self) -> &reqwest::Client {
pub fn provider_http(&self) -> &Client {
&self.provider_http
}
pub fn polling_http(&self) -> &reqwest::Client {
pub fn polling_http(&self) -> &Client {
&self.polling_http
}
@ -85,17 +85,18 @@ impl OcrClient {
}
#[cfg(any(test, feature = "test-support"))]
pub fn for_test(provider_http: reqwest::Client, document_http: reqwest::Client) -> Self {
pub fn for_test(provider_http: Client, no_redirect_http: Client) -> Self {
Self {
secrets: Arc::new(
litellm_secrets::source::EnvironmentSecrets::python_compatible(
provider_http.clone(),
),
),
provider_http,
polling_http: reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test polling client builds"),
document_fetcher: MediaFetcher::for_test(document_http),
polling_http: no_redirect_http.clone(),
document_fetcher: MediaFetcher::for_test(no_redirect_http),
vertex_auth: VertexAuth::default(),
settings: OcrSettings::default(),
secrets: Arc::new(litellm_secrets::source::EnvironmentSecrets::default()),
}
}
@ -311,7 +312,7 @@ mod tests {
let _connection = listener.accept().await.unwrap();
tokio::time::sleep(Duration::from_secs(1)).await;
});
let error = reqwest::Client::new()
let error = litellm_http::Client::plain_for_test()
.get(format!("http://{address}"))
.timeout(Duration::from_millis(10))
.send()

View file

@ -564,7 +564,10 @@ mod tests {
let params = ReductoParseV3Config
.map_ocr_params(&overrides, "parse-v3")
.unwrap();
let client = OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new());
let client = OcrClient::for_test(
litellm_http::Client::plain_for_test(),
litellm_http::Client::no_redirect_for_test(),
);
let connection = OcrConnection::default();
let document = serde_json::from_value(
json!({"type":"document_url","document_url":"reducto://ready.pdf"}),

View file

@ -19,7 +19,7 @@ async fn read_bounded(response: String, limit: usize) -> Result<bytes::Bytes, Er
socket.write_all(response.as_bytes()).await.unwrap();
std::future::pending::<()>().await;
});
let response = reqwest::Client::new()
let response = litellm_http::Client::plain_for_test()
.get(format!("http://{address}"))
.send()
.await

View file

@ -62,6 +62,7 @@ tokio = { workspace = true, features = ["rt", "sync"] }
url.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
litellm-secrets-aws.workspace = true
serde.workspace = true
serde_with.workspace = true

View file

@ -9,10 +9,10 @@ use super::{
cache_error,
config::{CacheBackendConfig, NativeCacheConfig, UnsupportedCacheConfig},
embedder::PythonEmbedder,
host_client,
native::NativeResponseCache,
};
use crate::errors::RustBridgeDeclined;
use crate::http::host_client;
fn declined(reason: UnsupportedCacheConfig) -> PyErr {
RustBridgeDeclined::new_err(reason.message())

View file

@ -1511,7 +1511,7 @@ mod tests {
path_service_account: Some("credentials.json".into()),
endpoint: litellm_cache_gcs::DEFAULT_ENDPOINT.into(),
},
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
Some("token".into()),
);
let matching_config = NativeCacheConfig {

View file

@ -1,3 +1,4 @@
use crate::http::host_client;
use crate::logger::run_sync_value;
use litellm_auth_aws::AwsAuthConfig;
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
@ -19,7 +20,6 @@ use super::{
config::{QdrantSemanticCacheConfig, project_redis_semantic},
embedder::PythonEmbedder,
facade::FacadeGuard,
host_client,
native::NativeResponseCache,
request::duration,
};

View file

@ -13,11 +13,9 @@ mod resolver;
mod semantic;
use litellm_cache::Error;
use litellm_http::ClientVariant;
use pyo3::{
exceptions::{PyNotImplementedError, PyRuntimeError, PyValueError},
prelude::*,
types::PyDict,
};
pub(crate) use self::{binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheResolver};
@ -29,11 +27,3 @@ fn cache_error(error: Error) -> PyErr {
_ => PyRuntimeError::new_err(error.to_string()),
}
}
/// The host's pooled HTTP client, configured from the proxy's HTTP settings.
fn host_client(py: Python<'_>, variant: ClientVariant) -> PyResult<reqwest::Client> {
let http_config = crate::http::call_config(py, &PyDict::new(py), true)?;
crate::http::pool()
.client(&http_config, variant)
.map_err(crate::http::client_error)
}

View file

@ -92,7 +92,7 @@ impl NativeResponseCache {
))
}
pub async fn s3(config: S3CacheConfig, http: reqwest::Client) -> Self {
pub async fn s3(config: S3CacheConfig, http: litellm_http::Client) -> Self {
let runtime = tokio::runtime::Handle::current();
let backend = S3Cache::new(config, http, ResponseCacheCodec, runtime);
let identity = BackendIdentity::S3 {
@ -112,7 +112,7 @@ impl NativeResponseCache {
Ok(Self::exact(ResponseCache::new(Arc::new(backend)), identity))
}
pub fn gcs(config: GcsConfig, client: reqwest::Client, token: Option<String>) -> Self {
pub fn gcs(config: GcsConfig, client: litellm_http::Client, token: Option<String>) -> Self {
let backend = match token {
Some(token) => GcsCache::with_token_source(
config,
@ -133,7 +133,7 @@ impl NativeResponseCache {
pub async fn azure_blob(
account_url: &str,
container: &str,
http: reqwest::Client,
http: litellm_http::Client,
) -> Result<Self, Error> {
let backend = AzureBlobCache::connect(
account_url,
@ -242,7 +242,7 @@ impl NativeResponseCache {
pub async fn qdrant_semantic(
config: QdrantSemanticCacheConfig,
client: reqwest::Client,
client: litellm_http::Client,
runtime: tokio::runtime::Handle,
) -> Result<Self, Error> {
let qdrant = qdrant_client::Qdrant::from_url(&config.grpc_url)

View file

@ -6,8 +6,8 @@ use std::{
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_http::{
HttpClientConfig, HttpClientPool, HttpSettings, HttpSettingsLayer, Resolution, SslVerify,
TlsSource, Unsupported,
Client, ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, HttpSettingsLayer,
Resolution, SslVerify, TlsSource, Unsupported,
media::{PublicDnsResolver, UrlPolicy},
};
use pyo3::{
@ -97,7 +97,10 @@ pub(crate) fn call_config(
let settings = HttpSettings::from_layers([
for_call(call_ssl_verify(kwargs)?, asynchronous),
HttpSettingsLayer::from_environment(&ProcessEnvironment),
configured(&PythonSettings::Http.read(py)?)?,
match PythonSettings::Http.read_or_unset(py)? {
Some(snapshot) => configured(&snapshot)?,
None => HttpSettingsLayer::default(),
},
])
.without_missing_files(&|path: &Path| path.exists());
let resolution = Resolution::from(&settings);
@ -107,6 +110,11 @@ pub(crate) fn call_config(
Ok(resolution.config)
}
pub(crate) fn host_client(py: Python<'_>, variant: ClientVariant) -> PyResult<Client> {
let config = call_config(py, &PyDict::new(py), true)?;
pool().client(&config, variant).map_err(client_error)
}
pub(crate) fn client_error(error: litellm_http::Error) -> PyErr {
match error {
litellm_http::Error::Read {
@ -143,7 +151,10 @@ fn unreported(
}
pub(crate) fn url_policy(py: Python<'_>) -> PyResult<UrlPolicy> {
project_url_policy(&PythonSettings::UrlPolicy.read(py)?)
match PythonSettings::UrlPolicy.read_or_unset(py)? {
Some(snapshot) => project_url_policy(&snapshot),
None => Ok(UrlPolicy::default()),
}
}
fn project_url_policy(snapshot: &Snapshot<'_>) -> PyResult<UrlPolicy> {

View file

@ -1,4 +1,4 @@
use pyo3::prelude::*;
use pyo3::{exceptions::PyModuleNotFoundError, prelude::*};
use crate::coercion::{FieldSpec, ProjectionError};
@ -40,6 +40,16 @@ impl PythonSettings {
Ok(Snapshot { group: self, value })
}
/// Reads the accessor, or `None` when the litellm package is not installed
/// (a bare extension module), meaning there are no configured values.
pub(crate) fn read_or_unset(self, py: Python<'_>) -> PyResult<Option<Snapshot<'_>>> {
match self.read(py) {
Ok(snapshot) => Ok(Some(snapshot)),
Err(error) if error.is_instance_of::<PyModuleNotFoundError>(py) => Ok(None),
Err(error) => Err(error),
}
}
#[cfg(test)]
pub(crate) fn snapshot(self, value: Bound<'_, PyAny>) -> Snapshot<'_> {
Snapshot { group: self, value }

View file

@ -3,7 +3,8 @@ use litellm_core::audio_transcription::{
Error, audio_transcription as run_audio_transcription, types::AudioTranscriptionRequest,
};
use litellm_host_python::from_py_argument;
use pyo3::prelude::*;
use litellm_http::HttpClientConfig;
use pyo3::{prelude::*, types::PyDict};
use serde_json::{Map, Value};
use crate::{
@ -12,6 +13,7 @@ use crate::{
};
async fn execute(
config: HttpClientConfig,
audio: Value,
optional_params: Map<String, Value>,
options: RouteOptions,
@ -24,16 +26,20 @@ async fn execute(
extra_headers,
timeout,
} = options;
run_audio_transcription(AudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
})
run_audio_transcription(
crate::http::pool(),
&config,
AudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
},
)
.await
}
@ -62,9 +68,10 @@ pub(crate) fn transcription(
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let config = crate::http::call_config(py, &PyDict::new(py), false)?;
run_sync(
py,
execute(audio, optional_params.unwrap_or_default(), options),
execute(config, audio, optional_params.unwrap_or_default(), options),
audio_transcription_error_to_pyerr,
)
}
@ -94,9 +101,10 @@ pub(crate) fn atranscription<'py>(
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let config = crate::http::call_config(py, &PyDict::new(py), true)?;
run_async(
py,
execute(audio, optional_params.unwrap_or_default(), options),
execute(config, audio, optional_params.unwrap_or_default(), options),
audio_transcription_error_to_pyerr,
)
}

View file

@ -7,6 +7,7 @@ use litellm_core::chat_completions::{
types::ChatCompletionsRequest,
};
use litellm_host_python::from_py_argument;
use litellm_http::HttpClientConfig;
use litellm_types::utils::ChatCompletionsResponse;
use pyo3::prelude::*;
use serde_json::{Map, Value};
@ -20,6 +21,7 @@ use crate::{
};
async fn execute(
config: HttpClientConfig,
messages: Vec<Value>,
optional_params: Map<String, Value>,
options: RouteOptions,
@ -32,16 +34,20 @@ async fn execute(
extra_headers,
timeout,
} = options;
run_chat_completions(ChatCompletionsRequest {
model: &model,
messages: Value::Array(messages),
optional_params,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
})
run_chat_completions(
crate::http::pool(),
&config,
ChatCompletionsRequest {
model: &model,
messages: Value::Array(messages),
optional_params,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
},
)
.await
}
@ -87,9 +93,15 @@ pub(crate) fn chat_completions(
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let config = crate::http::call_config(py, &PyDict::new(py), false)?;
run_sync(
py,
execute(messages, optional_params.unwrap_or_default(), options),
execute(
config,
messages,
optional_params.unwrap_or_default(),
options,
),
chat_completions_error_to_pyerr,
)
}
@ -119,9 +131,15 @@ pub(crate) fn achat_completions<'py>(
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let config = crate::http::call_config(py, &PyDict::new(py), true)?;
run_async(
py,
execute(messages, optional_params.unwrap_or_default(), options),
execute(
config,
messages,
optional_params.unwrap_or_default(),
options,
),
chat_completions_error_to_pyerr,
)
}

View file

@ -27,11 +27,14 @@ fn run_messages(
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let secrets = crate::secrets::source(py)?;
let config = crate::http::call_config(py, &kwargs, asynchronous)?;
let machine = messages_machine(crate::http::pool(), &config, secrets)
.map_err(crate::http::client_error)?;
run_legacy_call(
py,
SURFACE,
PublicCall::capture(&request, &args, &kwargs)?,
crate::logger::LoggedMachine::new(messages_machine(secrets)),
crate::logger::LoggedMachine::new(machine),
MessagesPythonHost::new(request.unbind()),
crate::preflight::sdk_preflight,
asynchronous,

View file

@ -210,7 +210,7 @@ handler.get_secret_from_manager = get_secret_from_manager
KeyManagementSettings::default(),
)),
Arc::new(move |_: &str| fallback.map(str::to_owned)),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
)
.with_failure_policy(FailurePolicy::EnvironmentFallback);
(resolver, locals, handler)

View file

@ -27,7 +27,12 @@ const NATIVE: FieldSpec<bool> = FieldSpec::new("native", |field| field.schema_bo
pub(crate) fn source(py: Python<'_>) -> PyResult<Arc<dyn SecretSource>> {
if PythonSettings::SecretManager.read(py)?.read(&NATIVE)? {
let context = litellm_host_python::PythonContext::capture(py)?;
return Ok(Arc::new(ResolvedSecrets::new(config::read(py)?, context)));
let client = crate::http::host_client(py, litellm_http::ClientVariant::Provider)?;
return Ok(Arc::new(ResolvedSecrets::new(
config::read(py)?,
context,
client,
)));
}
Ok(Arc::new(PythonSecrets::new(py)?))
}

View file

@ -3,6 +3,7 @@ use std::sync::Arc;
use futures_util::future::BoxFuture;
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_host_python::PythonContext;
use litellm_http::Client;
use litellm_secrets::source::SecretSource;
use litellm_secrets::{
Error, FailurePolicy, OidcResolver, SecretManagerState, SecretResolver, SecretValue,
@ -15,16 +16,20 @@ pub(crate) struct ResolvedSecrets {
}
impl ResolvedSecrets {
pub(crate) fn new(snapshot: SecretManagerSnapshot, context: PythonContext) -> Self {
Self::from_state(snapshot.into_state(context))
pub(crate) fn new(
snapshot: SecretManagerSnapshot,
context: PythonContext,
client: Client,
) -> Self {
Self::from_state(snapshot.into_state(context), client)
}
fn from_state(state: Arc<SecretManagerState>) -> Self {
fn from_state(state: Arc<SecretManagerState>, client: Client) -> Self {
Self {
resolver: SecretResolver::new_python_compatible(
state,
Arc::new(ProcessEnvironment),
OidcResolver::default(),
OidcResolver::new(client),
)
.with_failure_policy(FailurePolicy::EnvironmentFallback),
}
@ -79,7 +84,7 @@ mod tests {
}
async fn resolve(state: Arc<SecretManagerState>, name: &'static str) -> Option<String> {
ResolvedSecrets::from_state(state)
ResolvedSecrets::from_state(state, litellm_http::Client::plain_for_test())
.resolve(&[name])
.await
.unwrap()
@ -175,7 +180,10 @@ mod tests {
.expect(1)
.mount(&server)
.await;
let source = ResolvedSecrets::from_state(state(&server, KeyManagementSettings::default()));
let source = ResolvedSecrets::from_state(
state(&server, KeyManagementSettings::default()),
litellm_http::Client::plain_for_test(),
);
let snapshot = source.resolve(&[declared]).await.unwrap();
assert_eq!(snapshot.get(undeclared), None);
let result = source
@ -238,9 +246,12 @@ mod tests {
#[tokio::test]
async fn oidc_failures_are_not_converted_to_missing_secrets() {
let result = ResolvedSecrets::from_state(Arc::new(SecretManagerState::default()))
.resolve(&["oidc/"])
.await;
let result = ResolvedSecrets::from_state(
Arc::new(SecretManagerState::default()),
litellm_http::Client::plain_for_test(),
)
.resolve(&["oidc/"])
.await;
assert!(matches!(result, Err(litellm_secrets::Error::InvalidOidc)));
}

View file

@ -10,6 +10,7 @@ use litellm_secrets_types::PythonSecretRead;
use pyo3::{
exceptions::{PyAttributeError, PyRuntimeError, PyValueError},
prelude::*,
types::PyDict,
};
#[derive(Clone, PartialEq)]
@ -44,10 +45,18 @@ impl NativeSecretManager {
let system = configuration.system;
let settings = configuration.settings.clone();
let enterprise_enabled = configuration.enterprise_enabled;
let http_config = crate::http::call_config(py, &PyDict::new(py), false)?;
let backend = run_sync_value(py, async move {
load_native_manager(system, settings, environment, enterprise_enabled)
.await
.map_err(|error| PyValueError::new_err(error.to_string()))
load_native_manager(
crate::http::pool(),
&http_config,
system,
settings,
environment,
enterprise_enabled,
)
.await
.map_err(|error| PyValueError::new_err(error.to_string()))
})?;
Ok(Self {
backend,

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
litellm-http.workspace = true
tokio.workspace = true
litellm-auth-azure.workspace = true
litellm-auth-types.workspace = true
@ -18,6 +19,7 @@ veil.workspace = true
percent-encoding = "2.3"
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
wiremock = "0.6.5"
rstest.workspace = true
serde_json.workspace = true

View file

@ -19,7 +19,7 @@ const PATH_SEGMENT: &AsciiSet = &NON_ALPHANUMERIC
#[derive(Clone)]
pub struct AzureKeyVault {
client: reqwest::Client,
client: litellm_http::Client,
vault: reqwest::Url,
auth: Arc<AzureAuthService>,
inputs: Arc<AzureAuthInputs>,
@ -33,7 +33,7 @@ struct SecretResponse {
impl AzureKeyVault {
pub fn with_client(
client: reqwest::Client,
client: litellm_http::Client,
vault: reqwest::Url,
environment: Arc<dyn Lookup + Send + Sync>,
) -> Result<Self, Error> {
@ -57,7 +57,10 @@ impl AzureKeyVault {
})
}
pub fn new(environment: Arc<dyn Lookup + Send + Sync>) -> Result<Self, Error> {
pub fn new(
client: litellm_http::Client,
environment: Arc<dyn Lookup + Send + Sync>,
) -> Result<Self, Error> {
let value = environment
.get(AZURE_KEY_VAULT_URI)
.ok_or(Error::MissingEnvironment(AZURE_KEY_VAULT_URI))?;
@ -65,7 +68,7 @@ impl AzureKeyVault {
if vault.scheme() != "https" || vault.host_str().is_none() {
return Err(Error::VaultUri);
}
Self::with_client(reqwest::Client::new(), vault, environment)
Self::with_client(client, vault, environment)
}
pub fn scope(&self) -> &str {

View file

@ -11,7 +11,7 @@ use wiremock::{
fn manager(server: &MockServer) -> AzureKeyVault {
AzureKeyVault::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "fake".to_owned())),
)
@ -130,11 +130,14 @@ fn new_validates_vault_environment(
#[case] uri: Option<&'static str>,
#[case] missing_environment: bool,
) {
let result = AzureKeyVault::new(Arc::new(move |name: &str| {
(name == "AZURE_KEY_VAULT_URI")
.then(|| uri.map(str::to_owned))
.flatten()
}));
let result = AzureKeyVault::new(
litellm_http::Client::plain_for_test(),
Arc::new(move |name: &str| {
(name == "AZURE_KEY_VAULT_URI")
.then(|| uri.map(str::to_owned))
.flatten()
}),
);
if missing_environment {
assert!(matches!(
@ -155,7 +158,7 @@ fn new_validates_vault_environment(
#[case::local("http://localhost:8080", "https://localhost/.default")]
fn derives_scope_from_vault_host(#[case] uri: &str, #[case] expected: &str) {
let manager = AzureKeyVault::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
uri.parse().unwrap(),
Arc::new(|_: &str| None),
)
@ -184,7 +187,7 @@ async fn missing_credentials_do_not_request_vault() {
fn manager_without_credentials(server: &MockServer) -> AzureKeyVault {
AzureKeyVault::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
Arc::new(|name: &str| {
(name == "AZURE_CREDENTIAL").then(|| "ClientSecretCredential".to_owned())

View file

@ -10,7 +10,7 @@ use rstest::rstest;
#[ignore]
async fn reads_a_real_secret() {
let environment = Arc::new(ProcessEnvironment);
let manager = AzureKeyVault::new(environment).unwrap();
let manager = AzureKeyVault::new(litellm_http::Client::plain_for_test(), environment).unwrap();
let name = std::env::var("AZURE_KEY_VAULT_LIVE_SECRET_NAME").unwrap();
let secret = manager.get_secret(&name).await.unwrap().unwrap();
assert!(matches!(&secret, Secret::String(_)));

View file

@ -8,6 +8,7 @@ repository.workspace = true
[dependencies]
litellm-secrets-types.workspace = true
litellm-core-utils.workspace = true
litellm-http.workspace = true
base64.workspace = true
moka.workspace = true
reqwest.workspace = true
@ -19,6 +20,7 @@ percent-encoding = "2.3"
tokio = { workspace = true, features = ["sync"] }
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
rcgen = "0.14.10"
rstest.workspace = true
tempfile = "3.27.0"

View file

@ -18,6 +18,8 @@ pub enum Error {
MissingCredentials,
#[error("CyberArk client certificate could not be loaded")]
ClientCertificate,
#[error("CyberArk Conjur HTTP client could not be built")]
Client(#[redact] Box<litellm_http::Error>),
#[error("invalid refresh interval")]
RefreshInterval,
#[error("invalid CyberArk Conjur endpoint")]

View file

@ -2,10 +2,13 @@ mod client;
mod read;
mod write;
use std::{fs, sync::Arc, time::Duration};
use std::{sync::Arc, time::Duration};
use base64::{Engine, engine::general_purpose::STANDARD};
use litellm_core_utils::settings::Lookup;
use litellm_http::{
Client, ClientIdentity, ClientVariant, HttpClientConfig, HttpClientPool, TlsSource, Verify,
};
use litellm_secrets_types::{
BaseSecretManager, CyberarkOperationContext, RotationError, SecretCache, SecretDeleter,
SecretRotator, SecretValue, SecretWriteContext, SecretWriter, async_rotate_secret,
@ -37,7 +40,7 @@ const SECRET_NAME_SAFE: &AsciiSet = &NON_ALPHANUMERIC
#[derive(Clone)]
pub struct CyberArkSecretManager {
client: reqwest::Client,
client: Client,
endpoint: reqwest::Url,
account: String,
username: String,

View file

@ -2,7 +2,7 @@ use super::*;
impl CyberArkSecretManager {
pub fn with_client(
client: reqwest::Client,
client: Client,
endpoint: reqwest::Url,
account: String,
username: String,
@ -30,6 +30,8 @@ impl CyberArkSecretManager {
}
pub fn new(
pool: &HttpClientPool,
config: &HttpClientConfig,
environment: Arc<dyn Lookup + Send + Sync>,
enterprise_enabled: bool,
) -> Result<Self, Error> {
@ -46,21 +48,34 @@ impl CyberArkSecretManager {
.get(CYBERARK_SSL_VERIFY)
.map(|value| !value.trim().eq_ignore_ascii_case("false"))
.unwrap_or(true);
let mut builder = reqwest::Client::builder();
if !verify {
litellm_tracing::warn!(
"CyberArk SSL verification is disabled. This is insecure and should only be used for testing with self-signed certificates."
);
builder = builder.danger_accept_invalid_certs(true);
}
if !cert.is_empty() && !key.is_empty() {
let certificate = fs::read(cert).map_err(|_| Error::ClientCertificate)?;
let private_key = fs::read(key).map_err(|_| Error::ClientCertificate)?;
let identity = reqwest::Identity::from_pem(&[certificate, private_key].concat())
.map_err(|_| Error::ClientCertificate)?;
builder = builder.identity(identity);
}
let client = builder.build()?;
let config = HttpClientConfig {
verify: effective_verify(verify, &config.verify),
client_certificate: (!cert.is_empty() && !key.is_empty()).then(|| {
ClientIdentity::Split {
certificate: cert.into(),
key: key.into(),
}
}),
..config.clone()
};
let client =
pool.client(&config, ClientVariant::Provider)
.map_err(|error| match error {
litellm_http::Error::Read {
tls_source: TlsSource::ClientIdentity,
..
}
| litellm_http::Error::InvalidPem {
tls_source: TlsSource::ClientIdentity,
..
} => Error::ClientCertificate,
other => Error::Client(Box::new(other)),
})?;
let endpoint = reqwest::Url::parse(
&environment
.get(CYBERARK_API_BASE)
@ -139,9 +154,35 @@ impl CyberArkSecretManager {
}
}
fn effective_verify(cyberark_verify: bool, host: &Verify) -> Verify {
match (cyberark_verify, host) {
(false, _) => Verify::Disabled,
(true, Verify::Disabled) => Verify::BuiltInRoots,
(true, host) => host.clone(),
}
}
fn normalize_endpoint(mut endpoint: reqwest::Url) -> reqwest::Url {
if !endpoint.path().ends_with('/') {
endpoint.set_path(&format!("{}/", endpoint.path()));
}
endpoint
}
#[cfg(test)]
mod tests {
use std::path::PathBuf;
use super::*;
#[test]
fn cyberark_verification_does_not_follow_a_host_that_disabled_it() {
let bundle = Verify::CaBundle(PathBuf::from("/ca.pem"));
assert_eq!(
effective_verify(true, &Verify::Disabled),
Verify::BuiltInRoots
);
assert_eq!(effective_verify(true, &bundle), bundle);
assert_eq!(effective_verify(false, &bundle), Verify::Disabled);
}
}

View file

@ -7,6 +7,8 @@ use std::{
};
use base64::{Engine, engine::general_purpose::STANDARD};
use litellm_core_utils::settings::Lookup;
use litellm_http::{HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver};
use litellm_secrets_cyberark::{CyberArkSecretManager, DeleteOutcome, Error};
use litellm_secrets_types::{BaseSecretManager, CyberarkOperationContext, SecretValue};
use rstest::{fixture, rstest};

View file

@ -77,7 +77,7 @@ async fn authentication_encodes_login(#[case] username: &str, #[case] expected_p
.mount(&server)
.await;
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
"acct".into(),
username.into(),
@ -211,25 +211,25 @@ fn new_validates_credentials_before_license_and_configuration() {
let empty: Arc<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
Arc::new(|_: &str| None);
assert!(matches!(
CyberArkSecretManager::new(empty, true),
from_environment(empty, true),
Err(Error::MissingCredentials)
));
assert!(matches!(
CyberArkSecretManager::new(
from_environment(
Arc::new(|name: &str| (name == "CYBERARK_API_KEY").then(|| "k3y".into())),
false
),
Err(Error::EnterpriseRequired)
));
assert!(matches!(
CyberArkSecretManager::new(
from_environment(
Arc::new(|name: &str| (name == "CYBERARK_CLIENT_CERT").then(|| "cert".into())),
true
),
Err(Error::MissingCredentials)
));
assert!(matches!(
CyberArkSecretManager::new(
from_environment(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_REFRESH_INTERVAL" => Some("abc".into()),
@ -240,7 +240,7 @@ fn new_validates_credentials_before_license_and_configuration() {
Err(Error::RefreshInterval)
));
assert!(matches!(
CyberArkSecretManager::new(
from_environment(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_API_BASE" => Some("not a url".into()),
@ -254,7 +254,7 @@ fn new_validates_credentials_before_license_and_configuration() {
#[rstest]
fn certificate_only_credentials_are_validated_as_a_client_identity() {
let result = CyberArkSecretManager::new(
let result = from_environment(
Arc::new(|name: &str| match name {
"CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()),
"CYBERARK_CLIENT_KEY" => Some("/missing/key".into()),
@ -295,7 +295,7 @@ async fn configured_client_identity_preserves_auth_request_and_read_result(
let endpoint = server.uri();
let certificate = client_identity_directory.path().join("client.crt");
let key = client_identity_directory.path().join("client.key");
let manager = CyberArkSecretManager::new(
let manager = from_environment(
Arc::new(move |name: &str| match name {
"CYBERARK_API_BASE" => Some(endpoint.clone()),
"CYBERARK_API_KEY" => Some(api_key.into()),
@ -337,7 +337,7 @@ fn invalid_client_identity_is_not_ignored(
let certificate = client_identity_directory.path().join("client.crt");
let key = client_identity_directory.path().join("client.key");
let result = CyberArkSecretManager::new(
let result = from_environment(
Arc::new(move |name: &str| match name {
"CYBERARK_API_KEY" => Some(api_key.into()),
"CYBERARK_CLIENT_CERT" => Some(certificate.to_str().unwrap().into()),
@ -354,7 +354,7 @@ fn invalid_client_identity_is_not_ignored(
#[case::certificate_only("")]
#[case::certificate_and_api_key("k3y")]
fn client_identity_does_not_bypass_the_enterprise_requirement(#[case] api_key: &'static str) {
let result = CyberArkSecretManager::new(
let result = from_environment(
Arc::new(move |name: &str| match name {
"CYBERARK_API_KEY" => Some(api_key.into()),
"CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()),
@ -381,7 +381,7 @@ async fn new_reads_environment_defaults_end_to_end() {
.mount(&server)
.await;
let endpoint = server.uri();
let manager = CyberArkSecretManager::new(
let manager = from_environment(
Arc::new(move |name: &str| match name {
"CYBERARK_API_BASE" => Some(endpoint.clone()),
"CYBERARK_API_KEY" => Some("k3y".into()),
@ -404,7 +404,7 @@ async fn new_reads_environment_defaults_end_to_end() {
#[rstest]
fn new_reports_missing_client_certificate_files() {
assert!(matches!(
CyberArkSecretManager::new(
from_environment(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()),
@ -433,7 +433,7 @@ async fn trailing_slash_endpoint_preserves_base_path() {
.await;
let endpoint = format!("{}/prefix/", server.uri()).parse().unwrap();
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
endpoint,
"acct".into(),
"admin".into(),

View file

@ -48,9 +48,21 @@ pub(super) fn client_identity_directory() -> tempfile::TempDir {
directory
}
pub(super) fn from_environment(
environment: Arc<dyn Lookup + Send + Sync>,
enterprise_enabled: bool,
) -> Result<CyberArkSecretManager, Error> {
CyberArkSecretManager::new(
&HttpClientPool::new(Arc::new(PublicDnsResolver)),
&Resolution::from(&HttpSettings::default()).config,
environment,
enterprise_enabled,
)
}
pub(super) fn manager(server: &MockServer, ttl: Duration) -> CyberArkSecretManager {
CyberArkSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
"acct".into(),
"admin".into(),

View file

@ -166,7 +166,7 @@ async fn writes_match_python_parity_fixture(parity_fixture: ParityFixture) {
.mount(&server)
.await;
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
parity_fixture.account,
parity_fixture.username,
@ -230,7 +230,7 @@ async fn live_conjur_round_trip() {
.as_nanos()
);
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
endpoint.clone(),
account.clone(),
username.clone(),
@ -245,7 +245,7 @@ async fn live_conjur_round_trip() {
.await
.unwrap();
let verifier = CyberArkSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
endpoint.clone(),
account.clone(),
username.clone(),

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
litellm-http.workspace = true
moka.workspace = true
tokio.workspace = true
litellm-auth-gcp = { workspace = true, features = ["google-sdk"] }
@ -24,6 +25,7 @@ serde.workspace = true
reqwest.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
google-cloud-auth.workspace = true
rstest.workspace = true
wiremock = "0.6.5"

View file

@ -23,7 +23,7 @@ const CACHE_CAPACITY: u64 = 200;
#[derive(Clone)]
pub struct GoogleSecretManager {
client: reqwest::Client,
client: litellm_http::Client,
credentials: Arc<GoogleCredentials>,
endpoint: reqwest::Url,
project: String,
@ -46,7 +46,7 @@ struct Payload {
impl GoogleSecretManager {
pub fn with_client(
client: reqwest::Client,
client: litellm_http::Client,
endpoint: reqwest::Url,
project: String,
environment: Arc<dyn Lookup + Send + Sync>,
@ -79,6 +79,7 @@ impl GoogleSecretManager {
}
pub fn new(
client: litellm_http::Client,
environment: Arc<dyn Lookup + Send + Sync>,
enterprise_enabled: bool,
) -> Result<Self, Error> {
@ -104,7 +105,7 @@ impl GoogleSecretManager {
.get(GOOGLE_SECRET_MANAGER_ALWAYS_READ_SECRET_MANAGER)
.is_some_and(|v| v.eq_ignore_ascii_case("true"));
Self::with_client(
reqwest::Client::new(),
client,
reqwest::Url::parse("https://secretmanager.googleapis.com").expect("static URL"),
project,
environment,

View file

@ -11,7 +11,7 @@ use wiremock::{
fn manager(server: &MockServer, always_read: bool, ttl: Duration) -> GoogleSecretManager {
GoogleSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
"project".into(),
Arc::new(|name: &str| (name == "VERTEX_AI_API_KEY").then(|| "token".into())),
@ -214,11 +214,19 @@ async fn always_read_and_expired_cache_fetch_again(
#[rstest]
fn google_manager_requires_host_license_and_project_configuration() {
assert!(matches!(
GoogleSecretManager::new(Arc::new(|_: &str| None), false),
GoogleSecretManager::new(
litellm_http::Client::plain_for_test(),
Arc::new(|_: &str| None),
false
),
Err(Error::EnterpriseRequired)
));
assert!(matches!(
GoogleSecretManager::new(Arc::new(|_: &str| None), true),
GoogleSecretManager::new(
litellm_http::Client::plain_for_test(),
Arc::new(|_: &str| None),
true
),
Err(Error::MissingEnvironment(
"GOOGLE_SECRET_MANAGER_PROJECT_ID"
))
@ -236,7 +244,7 @@ fn google_manager_rejects_invalid_refresh_intervals(#[case] variable: &'static s
});
assert!(matches!(
GoogleSecretManager::new(environment, true),
GoogleSecretManager::new(litellm_http::Client::plain_for_test(), environment, true),
Err(Error::RefreshInterval)
));
}

View file

@ -23,6 +23,7 @@ litellm-secrets-hashicorp = { workspace = true, optional = true }
litellm-secrets-azure = { workspace = true, optional = true }
litellm-secrets-cyberark = { workspace = true, optional = true }
litellm-core-utils.workspace = true
litellm-http.workspace = true
base64.workspace = true
serde.workspace = true
strum.workspace = true
@ -33,6 +34,7 @@ moka.workspace = true
tokio = { workspace = true, features = ["fs"] }
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
rstest.workspace = true
wiremock = "0.6.5"
tempfile = "3"

View file

@ -10,6 +10,8 @@ pub enum Error {
InvalidCiphertext,
#[error("decrypted value is not UTF-8")]
Utf8,
#[error(transparent)]
Client(#[from] litellm_http::Error),
#[error("unsupported OIDC provider or missing build feature")]
UnsupportedOidc,
#[error("OIDC reference requires a provider and audience")]

View file

@ -1,10 +1,13 @@
use std::sync::Arc;
use litellm_core_utils::settings::Lookup;
use litellm_http::{HttpClientConfig, HttpClientPool};
use crate::{Error, KeyManagementSettings, KeyManagementSystem, SecretManager};
pub async fn load_native_manager(
pool: &HttpClientPool,
config: &HttpClientConfig,
system: KeyManagementSystem,
settings: KeyManagementSettings,
environment: Arc<dyn Lookup + Send + Sync>,
@ -29,14 +32,19 @@ pub async fn load_native_manager(
}
#[cfg(feature = "azure")]
(KeyManagementSystem::AzureKeyVault, _, environment, _) => Ok(
SecretManager::AzureKeyVault(crate::azure::AzureKeyVault::new(environment)?),
SecretManager::AzureKeyVault(crate::azure::AzureKeyVault::new(
pool.client(config, litellm_http::ClientVariant::Provider)?,
environment,
)?),
),
#[cfg(feature = "google")]
(KeyManagementSystem::GoogleSecretManager, _, environment, enterprise_enabled) => {
Ok(SecretManager::GoogleSecretManager(
crate::google::GoogleSecretManager::new(environment, enterprise_enabled)?,
))
}
(KeyManagementSystem::GoogleSecretManager, _, environment, enterprise_enabled) => Ok(
SecretManager::GoogleSecretManager(crate::google::GoogleSecretManager::new(
pool.client(config, litellm_http::ClientVariant::Provider)?,
environment,
enterprise_enabled,
)?),
),
#[cfg(feature = "google")]
(KeyManagementSystem::GoogleKms, _, environment, _) => {
crate::google::load_google_kms(Some(true), environment)
@ -51,11 +59,14 @@ pub async fn load_native_manager(
))
}
#[cfg(feature = "cyberark")]
(KeyManagementSystem::Cyberark, _, environment, enterprise_enabled) => {
Ok(SecretManager::Cyberark(
crate::cyberark::CyberArkSecretManager::new(environment, enterprise_enabled)?,
))
}
(KeyManagementSystem::Cyberark, _, environment, enterprise_enabled) => Ok(
SecretManager::Cyberark(crate::cyberark::CyberArkSecretManager::new(
pool,
config,
environment,
enterprise_enabled,
)?),
),
_ => Err(Error::NativeBackendUnavailable),
}
}

View file

@ -5,6 +5,7 @@ use std::{
use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
use litellm_core_utils::settings::Lookup;
use litellm_http::Client;
use moka::future::Cache;
use serde::Deserialize;
@ -82,7 +83,7 @@ impl NumericDate {
}
pub struct OidcResolver {
client: reqwest::Client,
client: Client,
google_identity_endpoint: reqwest::Url,
cache: Cache<String, (SecretValue, SystemTime)>,
clock: fn() -> SystemTime,
@ -90,25 +91,17 @@ pub struct OidcResolver {
azure_token_provider: std::sync::Arc<dyn litellm_secrets_azure::AzureTokenProvider>,
}
impl Default for OidcResolver {
fn default() -> Self {
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(600))
.connect_timeout(Duration::from_secs(5))
.build()
.expect("HTTP client configuration");
Self::new(
client,
reqwest::Url::parse("http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity").expect("static URL"),
)
}
}
const GOOGLE_IDENTITY_ENDPOINT: &str =
"http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity";
const REQUEST_TIMEOUT: Duration = Duration::from_secs(600);
impl OidcResolver {
pub fn new(client: reqwest::Client, google_identity_endpoint: reqwest::Url) -> Self {
pub fn new(client: Client) -> Self {
Self {
client,
google_identity_endpoint,
google_identity_endpoint: reqwest::Url::parse(GOOGLE_IDENTITY_ENDPOINT)
.expect("static URL"),
cache: Cache::builder()
.max_capacity(200)
.time_to_live(GOOGLE_TOKEN_MAX_TTL)
@ -121,6 +114,13 @@ impl OidcResolver {
}
}
pub fn with_google_identity_endpoint(self, google_identity_endpoint: reqwest::Url) -> Self {
Self {
google_identity_endpoint,
..self
}
}
#[cfg(feature = "azure")]
pub fn with_azure_token_provider(
self,
@ -180,6 +180,7 @@ impl OidcResolver {
let response = self
.client
.get(url)
.timeout(REQUEST_TIMEOUT)
.query(&[("audience", audience)])
.bearer_auth(authorization)
.header("Accept", "application/json; api-version=2.0")
@ -214,6 +215,7 @@ impl OidcResolver {
let response = self
.client
.get(self.google_identity_endpoint.clone())
.timeout(REQUEST_TIMEOUT)
.query(&[("audience", audience)])
.header("Metadata-Flavor", "Google")
.send()

View file

@ -1,10 +1,7 @@
use std::sync::Arc;
use crate::compatibility::python_manager_string;
use litellm_core_utils::{
serde_compat::parse_str_bool,
settings::{Lookup, ProcessEnvironment},
};
use litellm_core_utils::{serde_compat::parse_str_bool, settings::Lookup};
use crate::state::{LookupTarget, normalize_secret_name};
use crate::{Error, OidcResolver, Secret, SecretManagerState, SecretValue};
@ -24,16 +21,6 @@ pub struct SecretResolver {
python_compatible: bool,
}
impl Default for SecretResolver {
fn default() -> Self {
Self::new(
Arc::new(SecretManagerState::default()),
Arc::new(ProcessEnvironment),
OidcResolver::default(),
)
}
}
impl SecretResolver {
pub fn new(
state: Arc<SecretManagerState>,

View file

@ -37,15 +37,14 @@ impl SecretSource for SecretResolver {
}
}
#[derive(Default)]
pub struct EnvironmentSecrets(SecretResolver);
impl EnvironmentSecrets {
pub fn python_compatible() -> Self {
pub fn python_compatible(client: litellm_http::Client) -> Self {
Self(SecretResolver::new_python_compatible(
Arc::new(crate::SecretManagerState::default()),
Arc::new(litellm_core_utils::settings::ProcessEnvironment),
crate::OidcResolver::default(),
crate::OidcResolver::new(client),
))
}
}

View file

@ -53,7 +53,7 @@ async fn read_results_follow_the_selected_failure_policy(
},
)),
Arc::new(move |_: &str| environment.map(str::to_owned)),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
)
.with_failure_policy(policy);
let result = resolver
@ -97,7 +97,7 @@ async fn primary_secret_values_other_than_strings_resolve_to_none(
let resolver = SecretResolver::new_python_compatible(
Arc::new(state(&server, settings)),
Arc::new(|_: &str| Some("fallback".into())),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
);
let text = value.as_str();
assert_eq!(
@ -157,7 +157,7 @@ async fn gating_prediction_matches_actual_lookup(
let resolver = SecretResolver::new_python_compatible(
Arc::new(state),
Arc::new(|_: &str| Some("environment".into())),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
);
assert_eq!(
resolver

View file

@ -22,7 +22,7 @@ async fn azure_handler_reads_missing_and_failed_secrets() {
.await;
let manager = SecretManager::AzureKeyVault(
AzureKeyVault::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
std::sync::Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "fake".to_owned())),
)
@ -81,7 +81,7 @@ async fn successful_azure_responses_do_not_fall_back_when_the_value_is_empty_or_
.mount(&server)
.await;
let manager = AzureKeyVault::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "token".into())),
)
@ -92,7 +92,7 @@ async fn successful_azure_responses_do_not_fall_back_when_the_value_is_empty_or_
Default::default(),
)),
Arc::new(|_: &str| Some("environment".into())),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
);
assert_eq!(
resolver

View file

@ -59,7 +59,7 @@ fn manager(provider: Provider, server: &MockServer) -> SecretManager {
),
Provider::Azure => SecretManager::AzureKeyVault(
AzureKeyVault::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
environment,
)
@ -67,7 +67,7 @@ fn manager(provider: Provider, server: &MockServer) -> SecretManager {
),
Provider::Google => SecretManager::GoogleSecretManager(
GoogleSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
"project".into(),
environment,
@ -84,7 +84,7 @@ fn manager(provider: Provider, server: &MockServer) -> SecretManager {
.unwrap(),
),
Provider::Cyberark => SecretManager::Cyberark(CyberArkSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
"acct".into(),
"admin".into(),
@ -240,7 +240,7 @@ async fn python_read_failures_preserve_provider_fallback_rules(
KeyManagementSettings::default(),
)),
Arc::new(move |_: &str| environment_value.map(str::to_owned)),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
);
let expected = if matches!(provider, Provider::Aws) {
None

View file

@ -24,7 +24,7 @@ async fn cyberark_handler_reads_values_and_surfaces_errors() {
.mount(&server)
.await;
let manager = SecretManager::Cyberark(CyberArkSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
"acct".into(),
"admin".into(),

View file

@ -25,7 +25,7 @@ async fn google_resolver_distinguishes_absence_from_failure(#[case] status: u16)
_ => None,
});
let manager = GoogleSecretManager::with_client(
reqwest::Client::new(),
litellm_http::Client::plain_for_test(),
server.uri().parse().unwrap(),
"project".into(),
environment.clone(),
@ -40,7 +40,7 @@ async fn google_resolver_distinguishes_absence_from_failure(#[case] status: u16)
let resolver = SecretResolver::new_python_compatible(
Arc::new(state),
environment,
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
)
.with_failure_policy(FailurePolicy::Propagate);
let result = resolver.get_secret_str("KEY", None).await;

View file

@ -56,7 +56,7 @@ async fn hashicorp_handler_resolves_found_missing_and_failed_values() {
},
)),
Arc::new(|_: &str| None),
litellm_secrets::OidcResolver::default(),
litellm_secrets::OidcResolver::new(litellm_http::Client::plain_for_test()),
);
assert_eq!(
found_resolver
@ -131,7 +131,7 @@ async fn hashicorp_handler_resolves_found_missing_and_failed_values() {
let failed_resolver = SecretResolver::new_python_compatible(
Arc::new(failed_state),
Arc::new(|_: &str| None),
litellm_secrets::OidcResolver::default(),
litellm_secrets::OidcResolver::new(litellm_http::Client::plain_for_test()),
)
.with_failure_policy(FailurePolicy::Propagate);
assert!(matches!(

View file

@ -30,7 +30,7 @@ async fn environment_sources_resolve_expected_value(
("CIRCLE_OIDC_TOKEN_V2", "circle-v2"),
]);
assert_eq!(
OidcResolver::default()
OidcResolver::new(litellm_http::Client::plain_for_test())
.resolve(reference, env.as_ref())
.await
.unwrap()
@ -43,7 +43,7 @@ async fn environment_sources_resolve_expected_value(
#[tokio::test]
async fn environment_sources_bypass_boolean_conversion_and_defaults() {
let env = environment(&[("TOKEN", "true")]);
let oidc = OidcResolver::default();
let oidc = OidcResolver::new(litellm_http::Client::plain_for_test());
let resolver = SecretResolver::new(Arc::new(SecretManagerState::default()), env, oidc);
assert_eq!(
resolver
@ -94,7 +94,7 @@ async fn github_requests_are_authenticated_cached_and_revalidate_environment() {
),
("ACTIONS_ID_TOKEN_REQUEST_TOKEN", "request-token"),
]);
let oidc = OidcResolver::default();
let oidc = OidcResolver::new(litellm_http::Client::plain_for_test());
for _ in 0..2 {
assert_eq!(
oidc.resolve("oidc/github/https://service/oidc/path", env.as_ref())
@ -131,7 +131,7 @@ async fn file_allowlist_resolves_symlinks_while_environment_paths_remain_explici
("PATH_TOKEN", private.to_str().unwrap()),
("AZURE_FEDERATED_TOKEN_FILE", token.to_str().unwrap()),
]);
let oidc = OidcResolver::default();
let oidc = OidcResolver::new(litellm_http::Client::plain_for_test());
assert_eq!(
oidc.resolve(&format!("oidc/file/{}", token.display()), env.as_ref())
.await
@ -213,8 +213,9 @@ async fn google_expiry_caps_cache_and_preserves_audience(
.expect(calls)
.mount(&server)
.await;
let oidc =
OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()).with_clock(now);
let oidc = OidcResolver::new(litellm_http::Client::plain_for_test())
.with_google_identity_endpoint(server.uri().parse().unwrap())
.with_clock(now);
for _ in 0..2 {
assert_eq!(
oidc.resolve(
@ -234,7 +235,7 @@ async fn google_expiry_caps_cache_and_preserves_audience(
#[tokio::test]
async fn google_oidc_requires_its_build_feature() {
assert!(matches!(
OidcResolver::default()
OidcResolver::new(litellm_http::Client::plain_for_test())
.resolve("oidc/google/audience", environment(&[]).as_ref())
.await,
Err(Error::UnsupportedOidc)
@ -245,7 +246,7 @@ async fn google_oidc_requires_its_build_feature() {
#[tokio::test]
async fn azure_oidc_without_a_token_file_requires_its_build_feature() {
assert!(matches!(
OidcResolver::default()
OidcResolver::new(litellm_http::Client::plain_for_test())
.resolve("oidc/azure/scope", environment(&[]).as_ref())
.await,
Err(Error::UnsupportedOidc)
@ -261,7 +262,7 @@ async fn invalid_references_fail_before_environment_lookup(
#[case] reference: &str,
#[case] unsupported: bool,
) {
let error = OidcResolver::default()
let error = OidcResolver::new(litellm_http::Client::plain_for_test())
.resolve(reference, &|_: &str| {
panic!("invalid reference reached environment lookup")
})
@ -283,7 +284,8 @@ async fn unreadable_expiry_keeps_python_cache_fallback(#[case] token: &str) {
.expect(1)
.mount(&server)
.await;
let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap());
let resolver = OidcResolver::new(litellm_http::Client::plain_for_test())
.with_google_identity_endpoint(server.uri().parse().unwrap());
for _ in 0..2 {
assert_eq!(
resolver
@ -334,7 +336,8 @@ async fn azure_oidc_acquires_the_requested_scope_and_preserves_failures(#[case]
})
}
}
let oidc = OidcResolver::default().with_azure_token_provider(Arc::new(Provider(failed)));
let oidc = OidcResolver::new(litellm_http::Client::plain_for_test())
.with_azure_token_provider(Arc::new(Provider(failed)));
let resolver = SecretResolver::new_python_compatible(
Arc::new(SecretManagerState::default()),
environment(&[("AZURE_CLIENT_ID", "client-id")]),
@ -361,7 +364,7 @@ async fn azure_oidc_acquires_the_requested_scope_and_preserves_failures(#[case]
#[tokio::test]
async fn missing_oidc_environment_is_an_error(#[case] reference: &str) {
assert!(matches!(
OidcResolver::default()
OidcResolver::new(litellm_http::Client::plain_for_test())
.resolve(reference, environment(&[]).as_ref())
.await,
Err(Error::MissingEnvironment)
@ -380,7 +383,8 @@ async fn google_oidc_failures_are_not_cached_or_hidden_by_defaults() {
let resolver = SecretResolver::new_python_compatible(
Arc::new(SecretManagerState::default()),
environment(&[]),
OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()),
OidcResolver::new(litellm_http::Client::plain_for_test())
.with_google_identity_endpoint(server.uri().parse().unwrap()),
);
for _ in 0..2 {
assert!(matches!(
@ -420,7 +424,8 @@ async fn google_tokens_expire_at_the_python_cache_deadline(
.expect(2)
.mount(&server)
.await;
let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap())
let resolver = OidcResolver::new(litellm_http::Client::plain_for_test())
.with_google_identity_endpoint(server.uri().parse().unwrap())
.with_clock(|| UNIX_EPOCH + Duration::from_secs(1000));
assert_eq!(
resolver
@ -472,7 +477,8 @@ async fn google_cache_uses_payload_expiry_without_requiring_a_jwt_header() {
.expect(2)
.mount(&server)
.await;
let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap())
let resolver = OidcResolver::new(litellm_http::Client::plain_for_test())
.with_google_identity_endpoint(server.uri().parse().unwrap())
.with_clock(|| UNIX_EPOCH + Duration::from_secs(1000));
for _ in 0..2 {
assert_eq!(

View file

@ -12,7 +12,7 @@ fn resolver(value: Option<&str>) -> SecretResolver {
SecretResolver::new_python_compatible(
Arc::new(SecretManagerState::default()),
Arc::new(move |_: &str| value.clone()),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
)
}
@ -35,7 +35,7 @@ async fn native_reads_preserve_strings_and_report_conversion_errors(#[case] mana
let resolver = SecretResolver::new(
Arc::new(state),
Arc::new(move |_: &str| Some(raw.to_owned())),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
);
assert_eq!(
resolver
@ -75,7 +75,7 @@ async fn native_defaults_apply_to_absence_but_never_hide_provider_failures() {
KeyManagementSettings::default(),
)),
Arc::new(|_: &str| None),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
);
let result = resolver
.get_secret_str("key", Some(SecretValue::new("default")))
@ -189,7 +189,7 @@ fn managed(reply: Result<Option<Secret>, ()>, environment: Option<&'static str>)
KeyManagementSettings::default(),
)),
Arc::new(move |_: &str| environment.map(str::to_owned)),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
)
.with_failure_policy(FailurePolicy::EnvironmentFallback)
}
@ -281,7 +281,7 @@ async fn prefix_is_removed_once_and_resolved_from_environment() {
let resolver = SecretResolver::new_python_compatible(
Arc::new(state),
Arc::new(|name: &str| (name == "os.environ/KEY").then(|| "value".into())),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
);
assert_eq!(
resolver
@ -340,7 +340,7 @@ async fn excluded_hosted_keys_keep_the_python_manager_conversion_path(
},
)),
Arc::new(move |_: &str| Some(raw.to_owned())),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
);
assert_eq!(
resolver.get_secret("KEY", None).await.unwrap(),
@ -373,7 +373,7 @@ async fn azure_callback_absence_preserves_none_but_errors_fall_back(
KeyManagementSettings::default(),
)),
Arc::new(|_: &str| Some("environment".into())),
OidcResolver::default(),
OidcResolver::new(litellm_http::Client::plain_for_test()),
);
assert_eq!(
resolver

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