feat(rust): type the Vertex AI connection params and where Python reads them (#45641)

VertexParams in auth-types holds the vertex_* fields of CredentialLiteLLMParams
in both spellings, with one ParamSpec per setting: the wire names, the litellm
module global VertexBase.get_vertex_ai_* consults, and the environment names.
A credential deserializes from text or the JSON object a config spells, an
empty object counts as absent, and both credential fields are redacted in
Debug. auth-gcp is split so lib.rs is the entrypoint (config, auth,
constants), its lookups resolve through the specs, secret_names() derives
from them, and get_vertex_ai_project_from_credentials reads a service
account's project_id for hosts that need it before a token exchange. The OCR
host folds the module globals in by spec at projection, so OcrSettings no
longer carries vertex_project and vertex_location.
This commit is contained in:
yujonglee 2026-10-09 16:29:44 -07:00 • committed by GitHub
parent 0410abea8b
commit 04a95f60c1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
23 changed files with 1337 additions and 788 deletions

View file

@ -3411,9 +3411,12 @@ dependencies = [
"litellm-auth-types",
"moka",
"rstest",
"serde",
"serde_json",
"sha2 0.10.9",
"tempfile",
"tokio",
"veil",
]
[[package]]

View file

@ -1,18 +1,20 @@
pub const AWS_ACCESS_KEY_ID: &str = "AWS_ACCESS_KEY_ID";
pub const AWS_SECRET_ACCESS_KEY: &str = "AWS_SECRET_ACCESS_KEY";
pub const AWS_SESSION_TOKEN: &str = "AWS_SESSION_TOKEN";
pub const AWS_REGION_NAME: &str = "AWS_REGION_NAME";
pub const AWS_REGION: &str = "AWS_REGION";
use litellm_auth_types::AwsParams;
pub const AWS_ACCESS_KEY_ID: &str = AwsParams::ACCESS_KEY_ID.env[0];
pub const AWS_SECRET_ACCESS_KEY: &str = AwsParams::SECRET_ACCESS_KEY.env[0];
pub const AWS_SESSION_TOKEN: &str = AwsParams::SESSION_TOKEN.env[0];
pub const AWS_REGION_NAME: &str = AwsParams::REGION.env[0];
pub const AWS_REGION: &str = AwsParams::REGION.env[1];
pub const AWS_DEFAULT_REGION: &str = "AWS_DEFAULT_REGION";
pub const AWS_BEDROCK_RUNTIME_ENDPOINT: &str = "AWS_BEDROCK_RUNTIME_ENDPOINT";
pub const AWS_SESSION_NAME: &str = "AWS_SESSION_NAME";
pub const AWS_PROFILE_NAME: &str = "AWS_PROFILE_NAME";
pub const AWS_ROLE_NAME: &str = "AWS_ROLE_NAME";
pub const AWS_WEB_IDENTITY_TOKEN: &str = "AWS_WEB_IDENTITY_TOKEN";
pub const AWS_BEDROCK_RUNTIME_ENDPOINT: &str = AwsParams::BEDROCK_RUNTIME_ENDPOINT.env[0];
pub const AWS_SESSION_NAME: &str = AwsParams::SESSION_NAME.env[0];
pub const AWS_PROFILE_NAME: &str = AwsParams::PROFILE_NAME.env[0];
pub const AWS_ROLE_NAME: &str = AwsParams::ROLE_NAME.env[0];
pub const AWS_WEB_IDENTITY_TOKEN: &str = AwsParams::WEB_IDENTITY_TOKEN.env[0];
pub const AWS_ROLE_ARN: &str = "AWS_ROLE_ARN";
pub const AWS_WEB_IDENTITY_TOKEN_FILE: &str = "AWS_WEB_IDENTITY_TOKEN_FILE";
pub const AWS_STS_ENDPOINT: &str = "AWS_STS_ENDPOINT";
pub const AWS_EXTERNAL_ID: &str = "AWS_EXTERNAL_ID";
pub const AWS_STS_ENDPOINT: &str = AwsParams::STS_ENDPOINT.env[0];
pub const AWS_EXTERNAL_ID: &str = AwsParams::EXTERNAL_ID.env[0];
pub const AWS_BEARER_TOKEN_BEDROCK: &str = "AWS_BEARER_TOKEN_BEDROCK";
/// Headers SigV4 covers, beyond the `x-amz-` / `x-amzn-` prefixes. Mirrors
/// Python's `_filter_headers_for_aws_signature` allowlist.

View file

@ -12,9 +12,11 @@ google-sdk = ["dep:google-cloud-auth", "dep:http"]
litellm-auth-types.workspace = true
moka.workspace = true
serde.workspace = true
serde_json.workspace = true
sha2.workspace = true
tokio.workspace = true
veil.workspace = true
gcp_auth = "0.12.7"
google-cloud-auth = { workspace = true, optional = true }
@ -22,3 +24,4 @@ http = { workspace = true, optional = true }
[dev-dependencies]
rstest.workspace = true
tempfile.workspace = true

View file

@ -0,0 +1,274 @@
use std::{future::Future, path::Path, pin::Pin, sync::Arc};
use gcp_auth::{CustomServiceAccount, TokenProvider};
use litellm_auth_types::{CredentialPlacement, Error, ErrorDetail, http::apply_credential};
use moka::future::Cache;
use serde_json::Value;
use sha2::{Digest, Sha256};
use crate::{
config::{
CredentialSource, VertexConfig, credential_source, get_vertex_ai_project, non_empty_env,
},
constants::{
CLOUD_PLATFORM_SCOPE, GOOGLE_OAUTH_TOKEN_ENDPOINT, VERTEX_AI_API_KEY_ENV,
VERTEXAI_API_KEY_ENV,
},
};
pub struct VertexEnvironment {
pub headers: Vec<(String, String)>,
pub project_id: String,
}
struct VertexAccessToken {
token: String,
project_id: String,
}
#[derive(Clone)]
pub struct VertexAuth {
providers: Cache<CredentialCacheKey, Arc<dyn VertexTokenSource>>,
loader: Arc<dyn VertexProviderLoader>,
}
impl Default for VertexAuth {
fn default() -> Self {
Self::new(Arc::new(GcpProviderLoader))
}
}
impl VertexAuth {
pub fn new(loader: Arc<dyn VertexProviderLoader>) -> Self {
Self {
providers: Cache::builder().max_capacity(64).build(),
loader,
}
}
pub async fn access_token(
&self,
config: &VertexConfig,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<String, Error> {
self.load_provider(config, env_lookup).await?.token().await
}
pub async fn validate_environment(
&self,
headers: Vec<(String, String)>,
api_key: Option<&str>,
config: &VertexConfig,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<VertexEnvironment, Error> {
let has_authorization = headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("Authorization"));
let static_token = api_key
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
.or_else(|| non_empty_env(env_lookup, VERTEX_AI_API_KEY_ENV))
.or_else(|| non_empty_env(env_lookup, VERTEXAI_API_KEY_ENV));
let project_id = get_vertex_ai_project(config, env_lookup);
if !has_authorization && static_token.is_none() {
let access = self.get_access_token(config, env_lookup).await?;
return Ok(VertexEnvironment {
headers: apply_credential(headers, &access.token, CredentialPlacement::Bearer)?,
project_id: project_id.unwrap_or(access.project_id),
});
}
let project_id = match project_id {
Some(project_id) => project_id,
None => {
self.load_provider(config, env_lookup)
.await?
.project_id()
.await?
}
};
let headers = if has_authorization {
headers
} else {
apply_credential(
headers,
static_token.as_deref().expect("static token was checked"),
CredentialPlacement::Bearer,
)?
};
Ok(VertexEnvironment {
headers,
project_id,
})
}
async fn get_access_token(
&self,
config: &VertexConfig,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<VertexAccessToken, Error> {
let provider = self.load_provider(config, env_lookup).await?;
let (token, project_id) = tokio::try_join!(provider.token(), provider.project_id())?;
Ok(VertexAccessToken { token, project_id })
}
async fn load_provider(
&self,
config: &VertexConfig,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Arc<dyn VertexTokenSource>, Error> {
let source = credential_source(config, env_lookup);
let key = source.cache_key();
self.providers
.try_get_with(key, self.loader.load(source))
.await
.map_err(|error| (*error).clone())
}
}
pub trait VertexTokenSource: Send + Sync {
fn project_id(&self) -> VertexAuthFuture<'_, String>;
fn token(&self) -> VertexAuthFuture<'_, String>;
}
pub trait VertexProviderLoader: Send + Sync {
fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc<dyn VertexTokenSource>>;
}
pub type VertexAuthFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
struct GcpTokenSource(Arc<dyn TokenProvider>);
impl VertexTokenSource for GcpTokenSource {
fn project_id(&self) -> VertexAuthFuture<'_, String> {
Box::pin(async move {
self.0
.project_id()
.await
.map(|project| project.to_string())
.map_err(auth_acquisition_error)
})
}
fn token(&self) -> VertexAuthFuture<'_, String> {
Box::pin(async move {
self.0
.token(&[CLOUD_PLATFORM_SCOPE])
.await
.map(|token| token.as_str().to_string())
.map_err(auth_acquisition_error)
})
}
}
struct GcpProviderLoader;
impl VertexProviderLoader for GcpProviderLoader {
fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc<dyn VertexTokenSource>> {
Box::pin(async move {
let provider: Arc<dyn TokenProvider> = match source {
CredentialSource::Inline(configured) => Arc::new(
CustomServiceAccount::from_json(validate_request_credentials(
configured.expose(),
)?)
.map_err(auth_acquisition_error)?,
),
CredentialSource::Trusted(configured) => {
let configured = configured.expose();
let service_account = if Path::new(configured).is_file() {
CustomServiceAccount::from_file(configured)
} else {
CustomServiceAccount::from_json(configured)
}
.map_err(auth_acquisition_error)?;
Arc::new(service_account)
}
CredentialSource::ApplicationCredentials(path) => {
Arc::new(CustomServiceAccount::from_file(path).map_err(auth_acquisition_error)?)
}
CredentialSource::Adc => {
gcp_auth::provider().await.map_err(auth_acquisition_error)?
}
};
Ok(Arc::new(GcpTokenSource(provider)) as Arc<dyn VertexTokenSource>)
})
}
}
fn validate_request_credentials(configured: &str) -> Result<&str, Error> {
let token_uri = serde_json::from_str::<Value>(configured)
.ok()
.and_then(|credentials| {
credentials
.get("token_uri")
.and_then(Value::as_str)
.map(str::to_string)
});
if token_uri.as_deref() != Some(GOOGLE_OAUTH_TOKEN_ENDPOINT) {
return Err(Error::InvalidConfiguration(ErrorDetail::InvalidType {
field: "request-controlled vertex_credentials token_uri".into(),
expected: "the canonical Google OAuth token endpoint",
}));
}
Ok(configured)
}
impl CredentialSource {
fn cache_key(&self) -> CredentialCacheKey {
match self {
Self::Inline(configured) => {
CredentialCacheKey::Inline(Sha256::digest(configured.expose()).into())
}
Self::Trusted(configured) => {
CredentialCacheKey::Trusted(Sha256::digest(configured.expose()).into())
}
Self::ApplicationCredentials(path) => {
CredentialCacheKey::ApplicationCredentials(path.clone())
}
Self::Adc => CredentialCacheKey::Adc,
}
}
}
#[derive(Clone, Debug, Hash, PartialEq, Eq)]
enum CredentialCacheKey {
Inline([u8; 32]),
Trusted([u8; 32]),
ApplicationCredentials(String),
Adc,
}
fn auth_acquisition_error(error: gcp_auth::Error) -> Error {
Error::CredentialAcquisition(litellm_auth_types::ErrorDetail::failed(
"Vertex AI credentials",
error,
))
}
#[cfg(test)]
mod tests {
use litellm_auth_types::SecretValue;
use super::*;
#[rstest::rstest]
#[case::canonical_endpoint(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#, true)]
#[case::noncanonical_endpoint(r#"{"token_uri":"http://127.0.0.1/token"}"#, false)]
#[case::missing_endpoint("{}", false)]
fn request_credentials_require_canonical_token_endpoint(
#[case] credentials: &str,
#[case] accepted: bool,
) {
assert_eq!(validate_request_credentials(credentials).is_ok(), accepted);
}
#[rstest::rstest]
fn inline_and_trusted_credentials_never_share_a_cache_entry() {
assert_ne!(
CredentialSource::Inline(SecretValue::new("same-value")).cache_key(),
CredentialSource::Trusted(SecretValue::new("same-value")).cache_key()
);
}
}

View file

@ -0,0 +1,311 @@
use std::{collections::BTreeMap, path::Path};
use litellm_auth_types::{Error, ErrorDetail, InputSource, SecretValue, Sourced};
use serde_json::{Map, Value};
use litellm_auth_types::VertexParams;
use crate::constants::GOOGLE_APPLICATION_CREDENTIALS_ENV;
#[derive(Clone, Debug, Default)]
pub struct VertexConfig {
credentials: Option<Sourced<SecretValue>>,
project_id: Option<String>,
location: Option<String>,
}
impl VertexConfig {
pub fn new(
credentials: Option<Sourced<SecretValue>>,
project_id: Option<String>,
location: Option<String>,
) -> Self {
Self {
credentials: credentials.filter(|value| !value.value().expose().trim().is_empty()),
project_id: project_id.filter(|value| !value.trim().is_empty()),
location: location.filter(|value| !value.trim().is_empty()),
}
}
pub fn from_sourced_optional_params(
params: &Map<String, Value>,
sources: &BTreeMap<String, InputSource>,
) -> Result<Self, Error> {
Ok(Self::new(
optional_credentials(params, sources, VertexParams::CREDENTIALS.wire)?,
optional_string(params, VertexParams::PROJECT.wire)?,
optional_string(params, VertexParams::LOCATION.wire)?,
))
}
/// The deployment's `litellm_params`, which a host already trusts the way it trusts its
/// own configuration.
pub fn from_params(params: &VertexParams) -> Self {
Self::new(
params
.credentials()
.map(|value| Sourced::new(SecretValue::new(value), InputSource::Deployment)),
params.project(),
params.location(),
)
}
pub fn project_id(&self) -> Option<&str> {
self.project_id.as_deref()
}
pub fn location(&self) -> Option<&str> {
self.location.as_deref()
}
}
pub fn get_vertex_ai_project(
config: &VertexConfig,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Option<String> {
config
.project_id()
.map(str::to_string)
.or_else(|| VertexParams::PROJECT.resolve(&|_| None, env_lookup))
}
/// The project a request is billed to when nothing names it: the `project_id` inside the
/// service account the call would authenticate with. Application default credentials carry
/// no project that can be read without a token exchange, so they resolve to `None`.
pub fn get_vertex_ai_project_from_credentials(
config: &VertexConfig,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Option<String> {
let text = match credential_source(config, env_lookup) {
CredentialSource::Inline(configured) => configured.expose().to_string(),
CredentialSource::Trusted(configured) => {
let configured = configured.expose();
match std::fs::read_to_string(configured) {
Ok(contents) if Path::new(configured).is_file() => contents,
_ => configured.to_string(),
}
}
CredentialSource::ApplicationCredentials(path) => std::fs::read_to_string(path).ok()?,
CredentialSource::Adc => return None,
};
serde_json::from_str::<Value>(&text)
.ok()?
.get("project_id")
.and_then(Value::as_str)
.map(str::to_string)
}
pub fn get_vertex_ai_location(
config: &VertexConfig,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Option<String> {
config
.location()
.map(str::to_string)
.or_else(|| VertexParams::LOCATION.resolve(&|_| None, env_lookup))
}
#[derive(Clone, Debug)]
pub enum CredentialSource {
Inline(SecretValue),
Trusted(SecretValue),
ApplicationCredentials(String),
Adc,
}
pub(crate) fn credential_source(
config: &VertexConfig,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CredentialSource {
if let Some(configured) = config.credentials.clone() {
return match configured.source() {
InputSource::Request => CredentialSource::Inline(configured.into_value()),
InputSource::Deployment | InputSource::Environment => {
CredentialSource::Trusted(configured.into_value())
}
};
}
if let Some(configured) = VertexParams::CREDENTIALS.resolve(&|_| None, env_lookup) {
return CredentialSource::Trusted(SecretValue::new(configured));
}
non_empty_env(env_lookup, GOOGLE_APPLICATION_CREDENTIALS_ENV)
.map(CredentialSource::ApplicationCredentials)
.unwrap_or(CredentialSource::Adc)
}
fn optional_credentials(
params: &Map<String, Value>,
sources: &BTreeMap<String, InputSource>,
names: &[&str],
) -> Result<Option<Sourced<SecretValue>>, Error> {
for name in names {
let source = source_for(sources, name);
match params.get(*name) {
None | Some(Value::Null) => continue,
Some(Value::String(value)) if value.trim().is_empty() => continue,
Some(Value::String(value)) => {
return Ok(Some(Sourced::new(SecretValue::new(value), source)));
}
Some(Value::Object(value)) if value.is_empty() => continue,
Some(Value::Object(value)) => {
return serde_json::to_string(value)
.map(SecretValue::new)
.map(|value| Sourced::new(value, source))
.map(Some)
.map_err(|error| {
Error::InvalidConfiguration(ErrorDetail::failed(
"credential serialization",
error,
))
});
}
Some(_) => {
return Err(Error::InvalidConfiguration(ErrorDetail::InvalidType {
field: names[0].into(),
expected: "a string or null",
}));
}
}
}
Ok(None)
}
fn source_for(sources: &BTreeMap<String, InputSource>, name: &str) -> InputSource {
sources.get(name).copied().unwrap_or_default()
}
fn optional_string(params: &Map<String, Value>, names: &[&str]) -> Result<Option<String>, Error> {
for name in names {
match params.get(*name) {
None | Some(Value::Null) => continue,
Some(Value::String(value)) if value.trim().is_empty() => continue,
Some(Value::String(value)) => return Ok(Some(value.clone())),
Some(_) => {
return Err(Error::InvalidConfiguration(ErrorDetail::InvalidType {
field: names[0].into(),
expected: "a string or null",
}));
}
}
}
Ok(None)
}
pub(crate) fn non_empty_env(
env_lookup: &dyn Fn(&str) -> Option<String>,
name: &str,
) -> Option<String> {
env_lookup(name)
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
use crate::constants::VERTEXAI_CREDENTIALS_ENV;
fn config(value: Value) -> VertexConfig {
VertexConfig::from_sourced_optional_params(value.as_object().unwrap(), &BTreeMap::new())
.unwrap()
}
#[rstest::rstest]
fn empty_primary_values_fall_back_to_python_aliases() {
let config = config(json!({
"vertex_credentials": null,
"vertex_ai_credentials": "alias-credentials",
"vertex_project": " ",
"vertex_ai_project": "alias-project",
"vertex_location": null,
"vertex_ai_location": "alias-location"
}));
assert_eq!(
config.credentials.as_ref().unwrap().value().expose(),
"alias-credentials"
);
assert_eq!(config.project_id(), Some("alias-project"));
assert_eq!(config.location(), Some("alias-location"));
}
#[rstest::rstest]
fn typed_config_preserves_source_and_empty_value_fallback() {
let configured = VertexConfig::new(
Some(Sourced::new(
SecretValue::new("inline-json"),
InputSource::Request,
)),
Some("project".into()),
Some("location".into()),
);
assert!(matches!(
credential_source(&configured, &|_| Some("environment-json".into())),
CredentialSource::Inline(value) if value.expose() == "inline-json"
));
let empty = VertexConfig::new(
Some(Sourced::new(SecretValue::new(" "), InputSource::Request)),
Some(" ".into()),
Some(" ".into()),
);
assert!(matches!(
credential_source(&empty, &|_| None),
CredentialSource::Adc
));
assert_eq!(
get_vertex_ai_project(&empty, &|_| Some("env-project".into())).as_deref(),
Some("env-project")
);
assert_eq!(
get_vertex_ai_location(&empty, &|_| Some("env-location".into())).as_deref(),
Some("env-location")
);
}
#[rstest::rstest]
fn credential_discovery_prefers_input_then_environment_then_adc() {
let params = json!({"vertex_credentials":"input-json"});
let sources = BTreeMap::from([("vertex_credentials".to_string(), InputSource::Request)]);
let configured =
VertexConfig::from_sourced_optional_params(params.as_object().unwrap(), &sources)
.unwrap();
assert!(
matches!(credential_source(&configured, &|_| Some("environment-value".into())), CredentialSource::Inline(value) if value.expose() == "input-json")
);
let empty = VertexConfig::default();
assert!(
matches!(credential_source(&empty, &|name| (name == VERTEXAI_CREDENTIALS_ENV).then(|| "environment-json".into())), CredentialSource::Trusted(value) if value.expose() == "environment-json")
);
assert!(
matches!(credential_source(&empty, &|name| (name == GOOGLE_APPLICATION_CREDENTIALS_ENV).then(|| "adc.json".into())), CredentialSource::ApplicationCredentials(path) if path == "adc.json")
);
assert!(matches!(
credential_source(&empty, &|_| None),
CredentialSource::Adc
));
}
#[rstest::rstest]
fn params_become_a_deployment_sourced_config() {
let params = VertexParams {
vertex_credentials: Some("deployment-json".into()),
vertex_project: Some("project-1".into()),
vertex_ai_location: Some("europe-west4".into()),
..VertexParams::default()
};
let config = VertexConfig::from_params(&params);
assert_eq!(config.project_id(), Some("project-1"));
assert_eq!(config.location(), Some("europe-west4"));
assert!(
matches!(credential_source(&config, &|_| Some("environment-json".into())), CredentialSource::Trusted(value) if value.expose() == "deployment-json")
);
assert!(matches!(
credential_source(
&VertexConfig::from_params(&VertexParams::default()),
&|_| None
),
CredentialSource::Adc
));
}
}

View file

@ -0,0 +1,36 @@
use std::sync::LazyLock;
use litellm_auth_types::VertexParams;
pub const CLOUD_PLATFORM_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform";
pub const GOOGLE_OAUTH_TOKEN_ENDPOINT: &str = "https://oauth2.googleapis.com/token";
pub const GOOGLE_APPLICATION_CREDENTIALS_ENV: &str = "GOOGLE_APPLICATION_CREDENTIALS";
pub const VERTEX_AI_API_KEY_ENV: &str = "VERTEX_AI_API_KEY";
pub const VERTEXAI_API_KEY_ENV: &str = "VERTEXAI_API_KEY";
pub const VERTEXAI_CREDENTIALS_ENV: &str = VertexParams::CREDENTIALS.env[0];
pub const VERTEX_LOCATION_ENV: &str = VertexParams::LOCATION.env[1];
/// Every environment name the Vertex params and the token acquisition read, for hosts
/// that resolve secrets up front.
pub fn secret_names() -> &'static [&'static str] {
static NAMES: LazyLock<Vec<&'static str>> = LazyLock::new(|| {
[
VERTEX_AI_API_KEY_ENV,
VERTEXAI_API_KEY_ENV,
GOOGLE_APPLICATION_CREDENTIALS_ENV,
]
.into_iter()
.chain(
VertexParams::SPECS
.iter()
.flat_map(|spec| spec.env.iter().copied()),
)
.fold(Vec::new(), |mut names, name| {
if !names.contains(&name) {
names.push(name);
}
names
})
});
&NAMES
}

View file

@ -1,708 +1,16 @@
use std::{collections::BTreeMap, future::Future, path::Path, pin::Pin, sync::Arc};
use gcp_auth::{CustomServiceAccount, TokenProvider};
use litellm_auth_types::{
CredentialPlacement, Error, InputSource, SecretValue, Sourced, http::apply_credential,
};
use moka::future::Cache;
use serde_json::{Map, Value};
use sha2::{Digest, Sha256};
mod auth;
mod config;
pub mod constants;
#[cfg(feature = "google-sdk")]
mod sdk;
pub use auth::{
VertexAuth, VertexAuthFuture, VertexEnvironment, VertexProviderLoader, VertexTokenSource,
};
pub use config::{
CredentialSource, VertexConfig, get_vertex_ai_location, get_vertex_ai_project,
get_vertex_ai_project_from_credentials,
};
pub use constants::secret_names;
#[cfg(feature = "google-sdk")]
pub use sdk::GoogleCredentials;
const CLOUD_PLATFORM_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform";
const GOOGLE_OAUTH_TOKEN_ENDPOINT: &str = "https://oauth2.googleapis.com/token";
const GOOGLE_APPLICATION_CREDENTIALS_ENV: &str = "GOOGLE_APPLICATION_CREDENTIALS";
const VERTEX_AI_API_KEY_ENV: &str = "VERTEX_AI_API_KEY";
const VERTEXAI_API_KEY_ENV: &str = "VERTEXAI_API_KEY";
const VERTEXAI_CREDENTIALS_ENV: &str = "VERTEXAI_CREDENTIALS";
const VERTEXAI_PROJECT_ENV: &str = "VERTEXAI_PROJECT";
const VERTEXAI_LOCATION_ENV: &str = "VERTEXAI_LOCATION";
const VERTEX_LOCATION_ENV: &str = "VERTEX_LOCATION";
pub const SECRET_NAMES: &[&str] = &[
VERTEX_AI_API_KEY_ENV,
VERTEXAI_API_KEY_ENV,
VERTEXAI_CREDENTIALS_ENV,
GOOGLE_APPLICATION_CREDENTIALS_ENV,
VERTEXAI_PROJECT_ENV,
VERTEXAI_LOCATION_ENV,
VERTEX_LOCATION_ENV,
];
#[derive(Clone, Debug, Default)]
pub struct VertexConfig {
credentials: Option<Sourced<SecretValue>>,
project_id: Option<String>,
location: Option<String>,
}
impl VertexConfig {
pub fn new(
credentials: Option<Sourced<SecretValue>>,
project_id: Option<String>,
location: Option<String>,
) -> Self {
Self {
credentials: credentials.filter(|value| !value.value().expose().trim().is_empty()),
project_id: project_id.filter(|value| !value.trim().is_empty()),
location: location.filter(|value| !value.trim().is_empty()),
}
}
pub fn from_sourced_optional_params(
params: &Map<String, Value>,
sources: &BTreeMap<String, InputSource>,
) -> Result<Self, Error> {
Ok(Self::new(
optional_credentials(
params,
sources,
&["vertex_credentials", "vertex_ai_credentials"],
)?,
optional_string(params, &["vertex_project", "vertex_ai_project"])?,
optional_string(params, &["vertex_location", "vertex_ai_location"])?,
))
}
pub fn or_configured(self, project_id: Option<&str>, location: Option<&str>) -> Self {
let configured =
|value: Option<&str>| value.filter(|value| !value.is_empty()).map(str::to_string);
Self {
project_id: self.project_id.or_else(|| configured(project_id)),
location: self.location.or_else(|| configured(location)),
..self
}
}
pub fn project_id(&self) -> Option<&str> {
self.project_id.as_deref()
}
pub fn location(&self) -> Option<&str> {
self.location.as_deref()
}
}
pub struct VertexEnvironment {
pub headers: Vec<(String, String)>,
pub project_id: String,
}
struct VertexAccessToken {
token: String,
project_id: String,
}
pub fn get_vertex_ai_project(
config: &VertexConfig,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Option<String> {
config
.project_id()
.map(str::to_string)
.or_else(|| non_empty_env(env_lookup, VERTEXAI_PROJECT_ENV))
}
pub fn get_vertex_ai_location(
config: &VertexConfig,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Option<String> {
config
.location()
.map(str::to_string)
.or_else(|| non_empty_env(env_lookup, VERTEXAI_LOCATION_ENV))
.or_else(|| non_empty_env(env_lookup, VERTEX_LOCATION_ENV))
}
#[derive(Clone)]
pub struct VertexAuth {
providers: Cache<CredentialCacheKey, Arc<dyn VertexTokenSource>>,
loader: Arc<dyn VertexProviderLoader>,
}
impl Default for VertexAuth {
fn default() -> Self {
Self::new(Arc::new(GcpProviderLoader))
}
}
impl VertexAuth {
pub fn new(loader: Arc<dyn VertexProviderLoader>) -> Self {
Self {
providers: Cache::builder().max_capacity(64).build(),
loader,
}
}
pub async fn access_token(
&self,
config: &VertexConfig,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<String, Error> {
self.load_provider(config, env_lookup).await?.token().await
}
pub async fn validate_environment(
&self,
headers: Vec<(String, String)>,
api_key: Option<&str>,
config: &VertexConfig,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<VertexEnvironment, Error> {
let has_authorization = headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("Authorization"));
let static_token = api_key
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
.or_else(|| non_empty_env(env_lookup, VERTEX_AI_API_KEY_ENV))
.or_else(|| non_empty_env(env_lookup, VERTEXAI_API_KEY_ENV));
let project_id = get_vertex_ai_project(config, env_lookup);
if !has_authorization && static_token.is_none() {
let access = self.get_access_token(config, env_lookup).await?;
return Ok(VertexEnvironment {
headers: apply_credential(headers, &access.token, CredentialPlacement::Bearer)?,
project_id: project_id.unwrap_or(access.project_id),
});
}
let project_id = match project_id {
Some(project_id) => project_id,
None => {
self.load_provider(config, env_lookup)
.await?
.project_id()
.await?
}
};
let headers = if has_authorization {
headers
} else {
apply_credential(
headers,
static_token.as_deref().expect("static token was checked"),
CredentialPlacement::Bearer,
)?
};
Ok(VertexEnvironment {
headers,
project_id,
})
}
async fn get_access_token(
&self,
config: &VertexConfig,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<VertexAccessToken, Error> {
let provider = self.load_provider(config, env_lookup).await?;
let (token, project_id) = tokio::try_join!(provider.token(), provider.project_id())?;
Ok(VertexAccessToken { token, project_id })
}
async fn load_provider(
&self,
config: &VertexConfig,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Arc<dyn VertexTokenSource>, Error> {
let source = credential_source(config, env_lookup);
let key = source.cache_key();
self.providers
.try_get_with(key, self.loader.load(source))
.await
.map_err(|error| (*error).clone())
}
}
pub trait VertexTokenSource: Send + Sync {
fn project_id(&self) -> VertexAuthFuture<'_, String>;
fn token(&self) -> VertexAuthFuture<'_, String>;
}
pub trait VertexProviderLoader: Send + Sync {
fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc<dyn VertexTokenSource>>;
}
pub type VertexAuthFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
struct GcpTokenSource(Arc<dyn TokenProvider>);
impl VertexTokenSource for GcpTokenSource {
fn project_id(&self) -> VertexAuthFuture<'_, String> {
Box::pin(async move {
self.0
.project_id()
.await
.map(|project| project.to_string())
.map_err(auth_acquisition_error)
})
}
fn token(&self) -> VertexAuthFuture<'_, String> {
Box::pin(async move {
self.0
.token(&[CLOUD_PLATFORM_SCOPE])
.await
.map(|token| token.as_str().to_string())
.map_err(auth_acquisition_error)
})
}
}
struct GcpProviderLoader;
impl VertexProviderLoader for GcpProviderLoader {
fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc<dyn VertexTokenSource>> {
Box::pin(async move {
let provider: Arc<dyn TokenProvider> = match source {
CredentialSource::Inline(configured) => Arc::new(
CustomServiceAccount::from_json(validate_request_credentials(
configured.expose(),
)?)
.map_err(auth_acquisition_error)?,
),
CredentialSource::Trusted(configured) => {
let configured = configured.expose();
let service_account = if Path::new(configured).is_file() {
CustomServiceAccount::from_file(configured)
} else {
CustomServiceAccount::from_json(configured)
}
.map_err(auth_acquisition_error)?;
Arc::new(service_account)
}
CredentialSource::ApplicationCredentials(path) => {
Arc::new(CustomServiceAccount::from_file(path).map_err(auth_acquisition_error)?)
}
CredentialSource::Adc => {
gcp_auth::provider().await.map_err(auth_acquisition_error)?
}
};
Ok(Arc::new(GcpTokenSource(provider)) as Arc<dyn VertexTokenSource>)
})
}
}
fn validate_request_credentials(configured: &str) -> Result<&str, Error> {
let token_uri = serde_json::from_str::<Value>(configured)
.ok()
.and_then(|credentials| {
credentials
.get("token_uri")
.and_then(Value::as_str)
.map(str::to_string)
});
if token_uri.as_deref() != Some(GOOGLE_OAUTH_TOKEN_ENDPOINT) {
return Err(Error::InvalidConfiguration("request-controlled Vertex credentials must use the canonical Google OAuth token endpoint".into()));
}
Ok(configured)
}
#[derive(Clone, Debug)]
pub enum CredentialSource {
Inline(SecretValue),
Trusted(SecretValue),
ApplicationCredentials(String),
Adc,
}
impl CredentialSource {
fn cache_key(&self) -> CredentialCacheKey {
match self {
Self::Inline(configured) => {
CredentialCacheKey::Inline(Sha256::digest(configured.expose()).into())
}
Self::Trusted(configured) => {
CredentialCacheKey::Trusted(Sha256::digest(configured.expose()).into())
}
Self::ApplicationCredentials(path) => {
CredentialCacheKey::ApplicationCredentials(path.clone())
}
Self::Adc => CredentialCacheKey::Adc,
}
}
}
#[derive(Clone, Debug, Hash, PartialEq, Eq)]
enum CredentialCacheKey {
Inline([u8; 32]),
Trusted([u8; 32]),
ApplicationCredentials(String),
Adc,
}
fn credential_source(
config: &VertexConfig,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CredentialSource {
if let Some(configured) = config.credentials.clone() {
return match configured.source() {
InputSource::Request => CredentialSource::Inline(configured.into_value()),
InputSource::Deployment | InputSource::Environment => {
CredentialSource::Trusted(configured.into_value())
}
};
}
if let Some(configured) = non_empty_env(env_lookup, VERTEXAI_CREDENTIALS_ENV) {
return CredentialSource::Trusted(SecretValue::new(configured));
}
non_empty_env(env_lookup, GOOGLE_APPLICATION_CREDENTIALS_ENV)
.map(CredentialSource::ApplicationCredentials)
.unwrap_or(CredentialSource::Adc)
}
fn optional_credentials(
params: &Map<String, Value>,
sources: &BTreeMap<String, InputSource>,
names: &[&str],
) -> Result<Option<Sourced<SecretValue>>, Error> {
for name in names {
let source = source_for(sources, name);
match params.get(*name) {
None | Some(Value::Null) => continue,
Some(Value::String(value)) if value.trim().is_empty() => continue,
Some(Value::String(value)) => {
return Ok(Some(Sourced::new(SecretValue::new(value), source)));
}
Some(Value::Object(value)) if value.is_empty() => continue,
Some(Value::Object(value)) => {
return serde_json::to_string(value)
.map(SecretValue::new)
.map(|value| Sourced::new(value, source))
.map(Some)
.map_err(|error| {
Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed(
"credential serialization",
error,
))
});
}
Some(_) => {
return Err(Error::InvalidConfiguration(
litellm_auth_types::ErrorDetail::InvalidType {
field: names[0].into(),
expected: "a string or null",
},
));
}
}
}
Ok(None)
}
fn source_for(sources: &BTreeMap<String, InputSource>, name: &str) -> InputSource {
sources.get(name).copied().unwrap_or_default()
}
fn optional_string(params: &Map<String, Value>, names: &[&str]) -> Result<Option<String>, Error> {
for name in names {
match params.get(*name) {
None | Some(Value::Null) => continue,
Some(Value::String(value)) if value.trim().is_empty() => continue,
Some(Value::String(value)) => return Ok(Some(value.clone())),
Some(_) => {
return Err(Error::InvalidConfiguration(
litellm_auth_types::ErrorDetail::InvalidType {
field: names[0].into(),
expected: "a string or null",
},
));
}
}
}
Ok(None)
}
fn non_empty_env(env_lookup: &dyn Fn(&str) -> Option<String>, name: &str) -> Option<String> {
env_lookup(name)
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
fn auth_acquisition_error(error: gcp_auth::Error) -> Error {
Error::CredentialAcquisition(litellm_auth_types::ErrorDetail::failed(
"Vertex AI credentials",
error,
))
}
#[cfg(test)]
mod tests {
use std::collections::BTreeSet;
use std::sync::atomic::{AtomicUsize, Ordering};
use serde_json::json;
use super::*;
struct FakeProvider {
calls: Arc<AtomicUsize>,
}
impl VertexTokenSource for FakeProvider {
fn project_id(&self) -> VertexAuthFuture<'_, String> {
self.calls.fetch_add(1, Ordering::SeqCst);
Box::pin(async { Ok("adc-project".into()) })
}
fn token(&self) -> VertexAuthFuture<'_, String> {
self.calls.fetch_add(1, Ordering::SeqCst);
Box::pin(async { Ok("adc-token".into()) })
}
}
struct FakeLoader {
loads: Arc<AtomicUsize>,
provider: Arc<dyn VertexTokenSource>,
}
impl VertexProviderLoader for FakeLoader {
fn load(
&self,
_source: CredentialSource,
) -> VertexAuthFuture<'_, Arc<dyn VertexTokenSource>> {
let loads = self.loads.clone();
let provider = self.provider.clone();
Box::pin(async move {
loads.fetch_add(1, Ordering::SeqCst);
Ok(provider)
})
}
}
fn config(value: Value) -> VertexConfig {
VertexConfig::from_sourced_optional_params(value.as_object().unwrap(), &BTreeMap::new())
.unwrap()
}
fn auth(calls: Arc<AtomicUsize>, loads: Arc<AtomicUsize>) -> VertexAuth {
let provider: Arc<dyn VertexTokenSource> = Arc::new(FakeProvider { calls });
VertexAuth::new(Arc::new(FakeLoader { loads, provider }))
}
#[test]
fn config_is_typed_and_secrets_are_redacted() {
let config = config(json!({
"vertex_credentials":{"private_key":"secret-key"},
"vertex_project":"project-1",
"vertex_location":"europe-west4"
}));
assert_eq!(config.project_id(), Some("project-1"));
assert_eq!(config.location(), Some("europe-west4"));
assert!(!format!("{config:?}").contains("secret-key"));
assert!(
VertexConfig::from_sourced_optional_params(
json!({"vertex_credentials":true}).as_object().unwrap(),
&BTreeMap::new()
)
.is_err()
);
}
#[tokio::test]
async fn secret_names_cover_environment_reads() {
let seen = Arc::new(std::sync::Mutex::new(BTreeSet::<String>::new()));
let recorded = seen.clone();
let env = |name: &str| {
recorded.lock().unwrap().insert(name.to_string());
None
};
let auth = auth(Arc::new(AtomicUsize::new(0)), Arc::new(AtomicUsize::new(0)));
auth.validate_environment(Vec::new(), None, &VertexConfig::default(), &env)
.await
.unwrap();
get_vertex_ai_location(&VertexConfig::default(), &env);
assert!(
seen.lock()
.unwrap()
.iter()
.all(|name| SECRET_NAMES.contains(&name.as_str()))
);
}
#[test]
fn empty_primary_values_fall_back_to_python_aliases() {
let config = config(json!({
"vertex_credentials": null,
"vertex_ai_credentials": "alias-credentials",
"vertex_project": " ",
"vertex_ai_project": "alias-project",
"vertex_location": null,
"vertex_ai_location": "alias-location"
}));
assert_eq!(
config.credentials.as_ref().unwrap().value().expose(),
"alias-credentials"
);
assert_eq!(config.project_id(), Some("alias-project"));
assert_eq!(config.location(), Some("alias-location"));
}
#[test]
fn typed_config_preserves_source_and_empty_value_fallback() {
let configured = VertexConfig::new(
Some(Sourced::new(
SecretValue::new("inline-json"),
InputSource::Request,
)),
Some("project".into()),
Some("location".into()),
);
assert!(matches!(
credential_source(&configured, &|_| Some("environment-json".into())),
CredentialSource::Inline(value) if value.expose() == "inline-json"
));
let empty = VertexConfig::new(
Some(Sourced::new(SecretValue::new(" "), InputSource::Request)),
Some(" ".into()),
Some(" ".into()),
);
assert!(matches!(
credential_source(&empty, &|_| None),
CredentialSource::Adc
));
assert_eq!(
get_vertex_ai_project(&empty, &|_| Some("env-project".into())).as_deref(),
Some("env-project")
);
assert_eq!(
get_vertex_ai_location(&empty, &|_| Some("env-location".into())).as_deref(),
Some("env-location")
);
}
#[test]
fn project_and_location_prefer_input_then_environment() {
let configured =
config(json!({"vertex_project":"input-project","vertex_location":"input-location"}));
let env = |name: &str| Some(format!("env-{name}"));
assert_eq!(
get_vertex_ai_project(&configured, &env).as_deref(),
Some("input-project")
);
assert_eq!(
get_vertex_ai_location(&configured, &env).as_deref(),
Some("input-location")
);
let empty = VertexConfig::default();
assert_eq!(
get_vertex_ai_project(&empty, &|_| Some("env-project".into())).as_deref(),
Some("env-project")
);
assert_eq!(
get_vertex_ai_location(&empty, &|name| (name == VERTEX_LOCATION_ENV)
.then(|| "fallback-location".into()))
.as_deref(),
Some("fallback-location")
);
}
#[test]
fn credential_discovery_prefers_input_then_environment_then_adc() {
let params = json!({"vertex_credentials":"input-json"});
let sources = BTreeMap::from([("vertex_credentials".to_string(), InputSource::Request)]);
let configured =
VertexConfig::from_sourced_optional_params(params.as_object().unwrap(), &sources)
.unwrap();
assert!(
matches!(credential_source(&configured, &|_| Some("environment-value".into())), CredentialSource::Inline(value) if value.expose() == "input-json")
);
let empty = VertexConfig::default();
assert!(
matches!(credential_source(&empty, &|name| (name == VERTEXAI_CREDENTIALS_ENV).then(|| "environment-json".into())), CredentialSource::Trusted(value) if value.expose() == "environment-json")
);
assert!(
matches!(credential_source(&empty, &|name| (name == GOOGLE_APPLICATION_CREDENTIALS_ENV).then(|| "adc.json".into())), CredentialSource::ApplicationCredentials(path) if path == "adc.json")
);
assert!(matches!(
credential_source(&empty, &|_| None),
CredentialSource::Adc
));
assert_ne!(
CredentialSource::Inline(SecretValue::new("same-value")).cache_key(),
CredentialSource::Trusted(SecretValue::new("same-value")).cache_key()
);
}
#[rstest::rstest]
#[case::canonical_endpoint(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#, true)]
#[case::noncanonical_endpoint(r#"{"token_uri":"http://127.0.0.1/token"}"#, false)]
#[case::missing_endpoint("{}", false)]
fn request_credentials_require_canonical_token_endpoint(
#[case] credentials: &str,
#[case] accepted: bool,
) {
assert_eq!(validate_request_credentials(credentials).is_ok(), accepted);
}
#[tokio::test]
async fn explicit_token_and_header_do_not_acquire_adc() {
let loads = Arc::new(AtomicUsize::new(0));
let auth = auth(Arc::new(AtomicUsize::new(0)), loads.clone());
let configured = config(json!({"vertex_project":"project-1"}));
let explicit = auth
.validate_environment(Vec::new(), Some("access-token"), &configured, &|_| None)
.await
.unwrap();
assert_eq!(explicit.headers[0].1, "Bearer access-token");
let existing = auth
.validate_environment(
vec![("authorization".into(), "Bearer existing".into())],
None,
&configured,
&|_| None,
)
.await
.unwrap();
assert_eq!(existing.headers[0].1, "Bearer existing");
assert_eq!(loads.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn provider_is_reused_across_authentication_calls() {
let calls = Arc::new(AtomicUsize::new(0));
let loads = Arc::new(AtomicUsize::new(0));
let auth = auth(calls.clone(), loads.clone());
for _ in 0..2 {
let environment = auth
.validate_environment(Vec::new(), None, &VertexConfig::default(), &|_| None)
.await
.unwrap();
assert_eq!(environment.project_id, "adc-project");
assert_eq!(environment.headers[0].1, "Bearer adc-token");
}
assert_eq!(loads.load(Ordering::SeqCst), 1);
assert_eq!(calls.load(Ordering::SeqCst), 4);
}
#[test]
fn configured_defaults_sit_between_call_params_and_the_environment() {
let env = |name: &str| Some(format!("env-{name}"));
let from_config =
VertexConfig::default().or_configured(Some("global-project"), Some("global-location"));
assert_eq!(
get_vertex_ai_project(&from_config, &env).as_deref(),
Some("global-project")
);
assert_eq!(
get_vertex_ai_location(&from_config, &env).as_deref(),
Some("global-location")
);
let from_call =
config(json!({"vertex_project":"call-project","vertex_location":"call-location"}))
.or_configured(Some("global-project"), Some("global-location"));
assert_eq!(from_call.project_id(), Some("call-project"));
assert_eq!(from_call.location(), Some("call-location"));
let empty_global = VertexConfig::default().or_configured(Some(""), None);
assert_eq!(
get_vertex_ai_project(&empty_global, &env).as_deref(),
Some("env-VERTEXAI_PROJECT")
);
}
}

View file

@ -0,0 +1,120 @@
use std::{
collections::{BTreeMap, BTreeSet},
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
};
use litellm_auth_gcp::{
CredentialSource, VertexAuth, VertexAuthFuture, VertexConfig, VertexProviderLoader,
VertexTokenSource, get_vertex_ai_location, secret_names,
};
use rstest::rstest;
use serde_json::{Value, json};
struct FakeProvider {
calls: Arc<AtomicUsize>,
}
impl VertexTokenSource for FakeProvider {
fn project_id(&self) -> VertexAuthFuture<'_, String> {
self.calls.fetch_add(1, Ordering::SeqCst);
Box::pin(async { Ok("adc-project".into()) })
}
fn token(&self) -> VertexAuthFuture<'_, String> {
self.calls.fetch_add(1, Ordering::SeqCst);
Box::pin(async { Ok("adc-token".into()) })
}
}
struct FakeLoader {
loads: Arc<AtomicUsize>,
provider: Arc<dyn VertexTokenSource>,
}
impl VertexProviderLoader for FakeLoader {
fn load(&self, _source: CredentialSource) -> VertexAuthFuture<'_, Arc<dyn VertexTokenSource>> {
let loads = self.loads.clone();
let provider = self.provider.clone();
Box::pin(async move {
loads.fetch_add(1, Ordering::SeqCst);
Ok(provider)
})
}
}
fn config(value: Value) -> VertexConfig {
VertexConfig::from_sourced_optional_params(value.as_object().unwrap(), &BTreeMap::new())
.unwrap()
}
fn auth(calls: Arc<AtomicUsize>, loads: Arc<AtomicUsize>) -> VertexAuth {
let provider: Arc<dyn VertexTokenSource> = Arc::new(FakeProvider { calls });
VertexAuth::new(Arc::new(FakeLoader { loads, provider }))
}
#[rstest]
#[tokio::test]
async fn secret_names_cover_environment_reads() {
let seen = Arc::new(std::sync::Mutex::new(BTreeSet::<String>::new()));
let recorded = seen.clone();
let env = |name: &str| {
recorded.lock().unwrap().insert(name.to_string());
None
};
let auth = auth(Arc::new(AtomicUsize::new(0)), Arc::new(AtomicUsize::new(0)));
auth.validate_environment(Vec::new(), None, &VertexConfig::default(), &env)
.await
.unwrap();
get_vertex_ai_location(&VertexConfig::default(), &env);
assert!(
seen.lock()
.unwrap()
.iter()
.all(|name| secret_names().contains(&name.as_str()))
);
}
#[rstest]
#[tokio::test]
async fn explicit_token_and_header_do_not_acquire_adc() {
let loads = Arc::new(AtomicUsize::new(0));
let auth = auth(Arc::new(AtomicUsize::new(0)), loads.clone());
let configured = config(json!({"vertex_project":"project-1"}));
let explicit = auth
.validate_environment(Vec::new(), Some("access-token"), &configured, &|_| None)
.await
.unwrap();
assert_eq!(explicit.headers[0].1, "Bearer access-token");
let existing = auth
.validate_environment(
vec![("authorization".into(), "Bearer existing".into())],
None,
&configured,
&|_| None,
)
.await
.unwrap();
assert_eq!(existing.headers[0].1, "Bearer existing");
assert_eq!(loads.load(Ordering::SeqCst), 0);
}
#[rstest]
#[tokio::test]
async fn provider_is_reused_across_authentication_calls() {
let calls = Arc::new(AtomicUsize::new(0));
let loads = Arc::new(AtomicUsize::new(0));
let auth = auth(calls.clone(), loads.clone());
for _ in 0..2 {
let environment = auth
.validate_environment(Vec::new(), None, &VertexConfig::default(), &|_| None)
.await
.unwrap();
assert_eq!(environment.project_id, "adc-project");
assert_eq!(environment.headers[0].1, "Bearer adc-token");
}
assert_eq!(loads.load(Ordering::SeqCst), 1);
assert_eq!(calls.load(Ordering::SeqCst), 4);
}

View file

@ -0,0 +1,95 @@
use std::collections::BTreeMap;
use litellm_auth_gcp::{
VertexConfig,
constants::{GOOGLE_APPLICATION_CREDENTIALS_ENV, VERTEX_LOCATION_ENV},
get_vertex_ai_location, get_vertex_ai_project, get_vertex_ai_project_from_credentials,
};
use litellm_auth_types::VertexParams;
use rstest::rstest;
use serde_json::{Value, json};
fn config(value: Value) -> VertexConfig {
VertexConfig::from_sourced_optional_params(value.as_object().unwrap(), &BTreeMap::new())
.unwrap()
}
#[rstest]
fn config_is_typed_and_secrets_are_redacted() {
let config = config(json!({
"vertex_credentials":{"private_key":"secret-key"},
"vertex_project":"project-1",
"vertex_location":"europe-west4"
}));
assert_eq!(config.project_id(), Some("project-1"));
assert_eq!(config.location(), Some("europe-west4"));
assert!(!format!("{config:?}").contains("secret-key"));
assert!(
VertexConfig::from_sourced_optional_params(
json!({"vertex_credentials":true}).as_object().unwrap(),
&BTreeMap::new()
)
.is_err()
);
}
#[rstest]
fn project_and_location_prefer_input_then_environment() {
let configured =
config(json!({"vertex_project":"input-project","vertex_location":"input-location"}));
let env = |name: &str| Some(format!("env-{name}"));
assert_eq!(
get_vertex_ai_project(&configured, &env).as_deref(),
Some("input-project")
);
assert_eq!(
get_vertex_ai_location(&configured, &env).as_deref(),
Some("input-location")
);
let empty = VertexConfig::default();
assert_eq!(
get_vertex_ai_project(&empty, &|_| Some("env-project".into())).as_deref(),
Some("env-project")
);
assert_eq!(
get_vertex_ai_location(&empty, &|name| (name == VERTEX_LOCATION_ENV)
.then(|| "fallback-location".into()))
.as_deref(),
Some("fallback-location")
);
}
#[rstest]
fn project_is_read_from_the_credentials_that_would_authenticate() {
let key = tempfile::NamedTempFile::new().unwrap();
std::fs::write(key.path(), r#"{"project_id":"file-project"}"#).unwrap();
let path = key.path().to_str().unwrap().to_string();
let inline = VertexConfig::from_params(&VertexParams {
vertex_credentials: Some(r#"{"project_id":"inline-project"}"#.into()),
..VertexParams::default()
});
assert_eq!(
get_vertex_ai_project_from_credentials(&inline, &|_| None).as_deref(),
Some("inline-project")
);
let from_file = VertexConfig::from_params(&VertexParams {
vertex_credentials: Some(path.clone()),
..VertexParams::default()
});
assert_eq!(
get_vertex_ai_project_from_credentials(&from_file, &|_| None).as_deref(),
Some("file-project")
);
let empty = VertexConfig::default();
assert_eq!(
get_vertex_ai_project_from_credentials(&empty, &|name| (name
== GOOGLE_APPLICATION_CREDENTIALS_ENV)
.then(|| path.clone()))
.as_deref(),
Some("file-project")
);
assert_eq!(
get_vertex_ai_project_from_credentials(&empty, &|_| None),
None
);
}

View file

@ -53,7 +53,7 @@ pub use credential::{
};
pub use error::{Error, ErrorDetail, ErrorSource};
pub use http::CredentialPlacement;
pub use params::{AwsParams, ParamSpec};
pub use params::{AwsParams, ParamSpec, VertexParams};
pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy};
pub use secret::SecretValue;
pub use token::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};

View file

@ -1,6 +1,8 @@
mod aws;
mod vertex;
pub use aws::AwsParams;
pub use vertex::VertexParams;
/// One connection param of Python's `CredentialLiteLLMParams` and every place Python reads
/// it from, in order: the wire names on a call or a deployment, the `litellm.<name>` module

View file

@ -0,0 +1,118 @@
use serde::{Deserialize, Deserializer, Serialize, de::Unexpected};
use serde_json::Value;
use veil::Redact;
use super::ParamSpec;
/// The Vertex AI fields of Python's `GenericLiteLLMParams`, both the current names and the
/// `vertex_ai_*` spellings `VertexBase.safe_get_vertex_ai_*` still read.
#[derive(Redact, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct VertexParams {
#[serde(default, deserialize_with = "optional_credential_text")]
#[redact(with = "[REDACTED]")]
pub vertex_credentials: Option<String>,
#[serde(default)]
pub vertex_project: Option<String>,
#[serde(default)]
pub vertex_location: Option<String>,
#[serde(default, deserialize_with = "optional_credential_text")]
#[redact(with = "[REDACTED]")]
pub vertex_ai_credentials: Option<String>,
#[serde(default)]
pub vertex_ai_project: Option<String>,
#[serde(default)]
pub vertex_ai_location: Option<String>,
}
impl VertexParams {
pub const CREDENTIALS: ParamSpec = ParamSpec {
setting: "credentials",
wire: &["vertex_credentials", "vertex_ai_credentials"],
module_global: None,
env: &["VERTEXAI_CREDENTIALS"],
};
pub const PROJECT: ParamSpec = ParamSpec {
setting: "project",
wire: &["vertex_project", "vertex_ai_project"],
module_global: Some("vertex_project"),
env: &["VERTEXAI_PROJECT"],
};
pub const LOCATION: ParamSpec = ParamSpec {
setting: "location",
wire: &["vertex_location", "vertex_ai_location"],
module_global: Some("vertex_location"),
env: &["VERTEXAI_LOCATION", "VERTEX_LOCATION"],
};
/// Every Vertex AI param Python reads, with both spellings, the module global
/// `VertexBase.get_vertex_ai_*` consults, and the environment names it falls back to.
pub const SPECS: [ParamSpec; 3] = [Self::CREDENTIALS, Self::PROJECT, Self::LOCATION];
/// The wire names a host projects out of a caller's kwargs, derived from [`Self::SPECS`].
pub fn fields() -> impl Iterator<Item = &'static str> {
Self::SPECS
.iter()
.flat_map(|spec| spec.wire.iter().copied())
}
/// The value under a wire name, read through serde so the names can never drift from
/// the struct.
pub fn get(&self, wire: &str) -> Option<String> {
serde_json::to_value(self)
.ok()?
.get(wire)?
.as_str()
.map(str::to_string)
}
/// The spec's value from these params, then the environment.
pub fn resolve(
&self,
spec: &ParamSpec,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Option<String> {
spec.resolve(&|name| self.get(name), env_lookup)
}
pub fn credentials(&self) -> Option<String> {
self.resolve(&Self::CREDENTIALS, &|_| None)
}
pub fn project(&self) -> Option<String> {
self.resolve(&Self::PROJECT, &|_| None)
}
pub fn location(&self) -> Option<String> {
self.resolve(&Self::LOCATION, &|_| None)
}
}
/// A credential is a service account JSON text or a path to one; a YAML or kwargs mapping is
/// accepted as the JSON text it spells, the way Python passes a dict through.
fn optional_credential_text<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Option<String>, D::Error> {
match Option::<Value>::deserialize(deserializer)? {
None | Some(Value::Null) => Ok(None),
Some(Value::String(text)) => Ok(Some(text)),
Some(Value::Object(object)) if object.is_empty() => Ok(None),
Some(Value::Object(object)) => serde_json::to_string(&object)
.map(Some)
.map_err(serde::de::Error::custom),
Some(other) => Err(serde::de::Error::invalid_type(
Unexpected::Other(value_kind(&other)),
&"a string or an object",
)),
}
}
fn value_kind(value: &Value) -> &'static str {
match value {
Value::Null => "null",
Value::Bool(_) => "a boolean",
Value::Number(_) => "a number",
Value::String(_) => "a string",
Value::Array(_) => "a list",
Value::Object(_) => "an object",
}
}

View file

@ -0,0 +1,126 @@
use std::collections::BTreeSet;
use litellm_auth_types::VertexParams;
use rstest::rstest;
use serde_json::{Map, Value, json};
#[rstest]
#[case::current_spelling(
VertexParams { vertex_project: Some("p".into()), vertex_ai_project: Some("old".into()), ..VertexParams::default() },
Some("p")
)]
#[case::legacy_spelling(
VertexParams { vertex_ai_project: Some("old".into()), ..VertexParams::default() },
Some("old")
)]
#[case::blank_falls_through(
VertexParams { vertex_project: Some(" ".into()), vertex_ai_project: Some("old".into()), ..VertexParams::default() },
Some("old")
)]
#[case::unset(VertexParams::default(), None)]
fn params_prefer_the_current_spelling_over_the_legacy_one(
#[case] params: VertexParams,
#[case] project: Option<&str>,
) {
assert_eq!(params.project().as_deref(), project);
let mirrored = VertexParams {
vertex_credentials: params.vertex_project.clone(),
vertex_ai_credentials: params.vertex_ai_project.clone(),
vertex_location: params.vertex_project.clone(),
vertex_ai_location: params.vertex_ai_project.clone(),
..VertexParams::default()
};
assert_eq!(mirrored.credentials().as_deref(), project);
assert_eq!(mirrored.location().as_deref(), project);
}
#[rstest]
#[case::text(json!({"vertex_credentials": "/path/to/key.json"}), Some("/path/to/key.json"))]
#[case::object(json!({"vertex_credentials": {"type": "service_account"}}), Some(r#"{"type":"service_account"}"#))]
#[case::null(json!({"vertex_credentials": null}), None)]
#[case::empty_object_falls_through_to_the_legacy_spelling(
json!({"vertex_credentials": {}, "vertex_ai_credentials": "/legacy/key.json"}),
Some("/legacy/key.json")
)]
#[case::absent(json!({}), None)]
fn credentials_deserialize_from_text_or_an_object(
#[case] params: Value,
#[case] credentials: Option<&str>,
) {
let params: VertexParams = serde_json::from_value(params).unwrap();
assert_eq!(params.credentials().as_deref(), credentials);
}
#[rstest]
#[case::number(json!({"vertex_credentials": 7}))]
#[case::list(json!({"vertex_ai_credentials": ["a"]}))]
#[case::project(json!({"vertex_project": ["p"]}))]
fn a_param_of_the_wrong_type_is_rejected(#[case] params: Value) {
assert!(serde_json::from_value::<VertexParams>(params).is_err());
}
#[rstest]
fn fields_name_every_param_once() {
let filled: Value = VertexParams::fields()
.map(|name| (name.to_string(), Value::from(name)))
.collect::<Map<_, _>>()
.into();
let params: VertexParams = serde_json::from_value(filled).unwrap();
let expected = VertexParams {
vertex_credentials: Some("vertex_credentials".into()),
vertex_project: Some("vertex_project".into()),
vertex_location: Some("vertex_location".into()),
vertex_ai_credentials: Some("vertex_ai_credentials".into()),
vertex_ai_project: Some("vertex_ai_project".into()),
vertex_ai_location: Some("vertex_ai_location".into()),
};
assert_eq!(params, expected);
assert_eq!(
VertexParams::fields().collect::<BTreeSet<_>>().len(),
VertexParams::fields().count()
);
}
#[rstest]
fn debug_output_does_not_expose_credentials() {
let params = VertexParams {
vertex_credentials: Some("current-secret".into()),
vertex_ai_credentials: Some("legacy-secret".into()),
vertex_project: Some("project-1".into()),
..VertexParams::default()
};
let debug = format!("{params:?}");
assert!(!debug.contains("current-secret"));
assert!(!debug.contains("legacy-secret"));
assert!(debug.contains("project-1"));
}
#[rstest]
#[case::current_spelling(Some("p"), Some("legacy"), &[("VERTEXAI_PROJECT", "env")], Some("p"))]
#[case::legacy_spelling(None, Some("legacy"), &[("VERTEXAI_PROJECT", "env")], Some("legacy"))]
#[case::blank_params_fall_to_the_environment(Some(" "), None, &[("VERTEXAI_PROJECT", "env")], Some("env"))]
#[case::nothing(None, None, &[], None)]
fn a_spec_resolves_from_the_params_then_the_environment(
#[case] vertex_project: Option<&str>,
#[case] vertex_ai_project: Option<&str>,
#[case] environment: &[(&str, &str)],
#[case] expected: Option<&str>,
) {
let params = VertexParams {
vertex_project: vertex_project.map(str::to_string),
vertex_ai_project: vertex_ai_project.map(str::to_string),
..VertexParams::default()
};
let env = |name: &str| {
environment
.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| value.to_string())
};
assert_eq!(
params.resolve(&VertexParams::PROJECT, &env).as_deref(),
expected
);
}

View file

@ -1,6 +1,5 @@
use litellm_auth::{InputSource, Sourced};
use litellm_inference_ocr::arguments::is_supported_request;
use litellm_llms::base_llm::ocr::settings::OcrSettings;
use rstest::rstest;
use super::*;
@ -42,30 +41,6 @@ async fn mistral_is_served_at_the_resolved_project_and_location() {
);
}
#[tokio::test]
async fn configured_project_and_location_apply_when_the_call_sets_neither() {
let upstream = upstream([pages_response()]).await;
let route = ocr_route_with(OcrSettings {
vertex_project: Some("configured-project".into()),
vertex_location: Some("europe-west4".into()),
..OcrSettings::default()
});
route
.execute(
ocr_request("vertex_ai/mistral-ocr-maas", &upstream.uri(), json!({})),
&(),
None,
)
.await
.unwrap();
assert_eq!(
only_request(&upstream).await.url.path(),
"/v1/projects/configured-project/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
}
#[tokio::test]
async fn a_supplied_authorization_is_forwarded_without_a_static_token() {
let upstream = upstream([pages_response()]).await;

View file

@ -105,10 +105,7 @@ async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets(
owner,
&http,
Default::default(),
OcrSettings {
vertex_location: Some(location.into()),
..OcrSettings::default()
},
OcrSettings::default(),
Arc::new(RecordingSecrets::new([("VERTEXAI_CREDENTIALS", identity)])),
);
let request = decode_request(OcrWireRequest {
@ -118,7 +115,7 @@ async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets(
api_base: Some(upstream.uri()),
custom_llm_provider: None,
extra_headers: None,
optional_params: Default::default(),
optional_params: json!({"vertex_location": location}).as_object().cloned().unwrap(),
input_sources: Default::default(),
timeout_seconds: Some(5.0),
}).unwrap();

View file

@ -9,8 +9,6 @@ pub struct OcrSettings {
pub poll_timeout: Duration,
pub document_intelligence_api_version: String,
pub document_intelligence_dpi: i64,
pub vertex_project: Option<String>,
pub vertex_location: Option<String>,
pub enable_azure_ad_token_refresh: bool,
}
@ -22,8 +20,6 @@ impl Default for OcrSettings {
poll_timeout: Duration::from_secs(120),
document_intelligence_api_version: "2024-11-30".into(),
document_intelligence_dpi: 96,
vertex_project: None,
vertex_location: None,
enable_azure_ad_token_refresh: false,
}
}

View file

@ -7,15 +7,10 @@ use crate::base_llm::ocr::{
};
pub(super) fn vertex_config(request: &PreparedOcrRequest) -> Result<VertexConfig, Error> {
let settings = &request.connection.settings;
Ok(VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?
.or_configured(
settings.vertex_project.as_deref(),
settings.vertex_location.as_deref(),
))
)?)
}
pub(super) fn validate_destination(connection: &OcrConnection) -> Result<(), Error> {

View file

@ -35,7 +35,7 @@ impl BaseOcrConfig for VertexAiOcrConfig {
}
fn secret_names(&self) -> Vec<&'static str> {
litellm_auth_gcp::SECRET_NAMES.to_vec()
litellm_auth_gcp::secret_names().to_vec()
}
fn map_ocr_params(

View file

@ -68,7 +68,7 @@ fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult<boo
.value(py)
.getattr("name")?
.extract::<Option<String>>()?
.is_some_and(|name| name == expected))
.is_some_and(|name| name == expected || name.starts_with(&format!("{expected}."))))
}
#[cfg(test)]

View file

@ -13,7 +13,7 @@ use pyo3::{
use super::{
errors::to_pyerr as ocr_error_to_pyerr,
project::{OcrHostHandles, project_request},
project::{OcrHostHandles, SETTINGS_ERROR_MARKER, project_request},
};
use crate::marshal::public_response;
@ -68,7 +68,7 @@ impl OcrPythonHost {
}
fn map_failure(&self, py: Python<'_>, error: PyErr) -> PyErr {
if !error.is_instance_of::<PyException>(py) {
if !error.is_instance_of::<PyException>(py) || is_settings_error(py, &error) {
return error;
}
let provider = match &self.data {
@ -87,6 +87,15 @@ impl OcrPythonHost {
}
}
fn is_settings_error(py: Python<'_>, error: &PyErr) -> bool {
error
.value(py)
.getattr_opt(SETTINGS_ERROR_MARKER)
.ok()
.flatten()
.is_some()
}
impl PythonBinding for OcrPythonHost {
type Protocol = Ocr;
type Failure = PyErr;

View file

@ -19,10 +19,6 @@ use crate::{
python_settings::{PythonSettings, Snapshot},
};
const VERTEX_PROJECT: FieldSpec<Option<String>> =
FieldSpec::new("vertex_project", |field| field.falsy_optional_string());
const VERTEX_LOCATION: FieldSpec<Option<String>> =
FieldSpec::new("vertex_location", |field| field.falsy_optional_string());
const ENABLE_AZURE_AD_TOKEN_REFRESH: FieldSpec<bool> =
FieldSpec::new("enable_azure_ad_token_refresh", |field| {
Ok(field.exact_true())
@ -60,8 +56,6 @@ fn ocr_settings(py: Python<'_>) -> PyResult<OcrSettings> {
fn project_provider_defaults(snapshot: &Snapshot<'_>) -> PyResult<OcrSettings> {
Ok(OcrSettings {
vertex_project: snapshot.read(&VERTEX_PROJECT)?,
vertex_location: snapshot.read(&VERTEX_LOCATION)?,
enable_azure_ad_token_refresh: snapshot.read(&ENABLE_AZURE_AD_TOKEN_REFRESH)?,
..OcrSettings::from_environment(&ProcessEnvironment)
})
@ -108,31 +102,29 @@ mod tests {
use crate::python_settings::PythonSettings;
#[rstest::rstest]
fn provider_defaults_distinguish_falsey_values_and_exact_true() {
fn provider_defaults_read_exact_true_only() {
Python::initialize();
Python::attach(|py| {
let value = py.eval(c"__import__('types').SimpleNamespace(vertex_project=[], vertex_location=0, enable_azure_ad_token_refresh=1)", None, None).unwrap();
let value = py
.eval(
c"__import__('types').SimpleNamespace(enable_azure_ad_token_refresh=1)",
None,
None,
)
.unwrap();
let snapshot = PythonSettings::ProviderDefaults.snapshot(value.clone());
let projected = super::project_provider_defaults(&snapshot).unwrap();
assert_eq!(projected.vertex_project, None);
assert_eq!(projected.vertex_location, None);
assert!(!projected.enable_azure_ad_token_refresh);
value.setattr("vertex_project", "project").unwrap();
value.setattr("vertex_location", "region").unwrap();
assert!(
!super::project_provider_defaults(&snapshot)
.unwrap()
.enable_azure_ad_token_refresh
);
value
.setattr("enable_azure_ad_token_refresh", true)
.unwrap();
let next = super::project_provider_defaults(&snapshot).unwrap();
assert_eq!(next.vertex_project.as_deref(), Some("project"));
assert_eq!(next.vertex_location.as_deref(), Some("region"));
assert!(next.enable_azure_ad_token_refresh);
value.setattr("vertex_project", 1).unwrap();
let error = super::project_provider_defaults(&snapshot).err().unwrap();
assert!(error.is_instance_of::<pyo3::exceptions::PyValueError>(py));
assert!(
error
.to_string()
.contains("provider_defaults.vertex_project")
super::project_provider_defaults(&snapshot)
.unwrap()
.enable_azure_ad_token_refresh
);
});
}

View file

@ -1,4 +1,5 @@
use litellm_auth::SecretValue;
use litellm_auth::VertexParams;
use litellm_host_python::{from_py, present};
use litellm_inference_ocr::{
types::{LiteLLMOcrRequest, OcrDocumentInput},
@ -10,8 +11,10 @@ use serde_json::{Map, Value};
use super::{document::FileDocumentInput, errors::to_pyerr as ocr_error_to_pyerr};
use crate::{
coercion::FieldSpec,
credentials::{self, CallerTokenProvider},
marshal::{project_optional_fields, python_timeout_seconds, request_input_sources},
python_settings::{PythonSettings, Snapshot},
};
/// What the host keeps after projection: the caller's token callable that answers the
@ -123,6 +126,16 @@ pub(super) fn project_request(
let names = specs.iter().map(|spec| spec.name).collect::<Vec<_>>();
let optional_params =
project_optional_fields(names.iter().copied(), |name| present(kwargs, bound, name))?;
let globals = module_globals_to_read(&optional_params, &names);
let defaults = if globals.is_empty() {
None
} else {
PythonSettings::ProviderDefaults.read_or_unset(bound.py())?
};
let optional_params = match defaults {
Some(defaults) => fold_module_globals(bound.py(), optional_params, &globals, &defaults)?,
None => optional_params,
};
let input_sources = request_input_sources(
kwargs,
names
@ -156,6 +169,59 @@ pub(super) fn project_request(
))
}
/// Python's `VertexBase.get_vertex_ai_project` and `get_vertex_ai_location` fall back to a
/// `litellm.<name>` module global when a call names neither spelling: the globals of the
/// specs this route consumes whose wire names the call left out.
fn module_globals_to_read(
optional_params: &Map<String, Value>,
consumed: &[&str],
) -> Vec<&'static str> {
VertexParams::SPECS
.iter()
.filter(|spec| spec.wire.iter().all(|name| consumed.contains(name)))
.filter(|spec| {
!spec
.wire
.iter()
.any(|name| optional_params.contains_key(*name))
})
.filter_map(|spec| spec.module_global)
.collect()
}
/// Folds each module global in under the name its spec declares, skipping falsy values the
/// way Python's `or` chain does.
fn fold_module_globals(
py: Python<'_>,
optional_params: Map<String, Value>,
globals: &[&'static str],
defaults: &Snapshot<'_>,
) -> PyResult<Map<String, Value>> {
let read = globals
.iter()
.map(|name| {
let spec = FieldSpec::new(name, |field| field.falsy_optional_string());
Ok(defaults
.read(&spec)
.map_err(|error| settings_error(py, error.into()))?
.map(|value| (name.to_string(), Value::String(value))))
})
.collect::<PyResult<Vec<_>>>()?;
Ok(optional_params
.into_iter()
.chain(read.into_iter().flatten())
.collect())
}
/// A misconfigured litellm setting is the operator's error, not the request's or the
/// provider's, so the host raises it as is instead of mapping it onto a provider failure.
pub(super) const SETTINGS_ERROR_MARKER: &str = "native_settings_error";
fn settings_error(py: Python<'_>, error: PyErr) -> PyErr {
error.value(py).setattr(SETTINGS_ERROR_MARKER, true).ok();
error
}
#[cfg(test)]
mod tests {
use litellm_llms_types::formats::ocr::OcrDocument;
@ -163,6 +229,105 @@ mod tests {
use super::*;
#[rstest::rstest]
#[case::global_fills_a_missing_project(
serde_json::json!({"vertex_location": "us-east5"}),
"vertex_project='from-global', vertex_location='ignored'",
serde_json::json!({"vertex_location": "us-east5", "vertex_project": "from-global"}),
)]
#[case::either_spelling_on_the_call_wins(
serde_json::json!({"vertex_ai_project": "from-call"}),
"vertex_project='from-global', vertex_location=None",
serde_json::json!({"vertex_ai_project": "from-call"}),
)]
#[case::falsy_globals_are_absent(
serde_json::json!({}),
"vertex_project=[], vertex_location=''",
serde_json::json!({}),
)]
fn module_globals_fill_the_vertex_params_the_call_leaves_out(
#[case] params: Value,
#[case] defaults: &str,
#[case] expected: Value,
) {
Python::initialize();
Python::attach(|py| {
let namespace = py
.eval(
&std::ffi::CString::new(format!(
"__import__('types').SimpleNamespace({defaults})"
))
.unwrap(),
None,
None,
)
.unwrap();
let snapshot = PythonSettings::ProviderDefaults.snapshot(namespace);
let consumed = VertexParams::fields().collect::<Vec<_>>();
let params = params.as_object().cloned().unwrap();
let globals = module_globals_to_read(&params, &consumed);
let folded = fold_module_globals(py, params, &globals, &snapshot).unwrap();
assert_eq!(Value::Object(folded), expected);
});
}
#[rstest::rstest]
fn a_bad_global_is_a_marked_settings_error() {
Python::initialize();
Python::attach(|py| {
let namespace = py
.eval(
c"__import__('types').SimpleNamespace(vertex_project=1)",
None,
None,
)
.unwrap();
let snapshot = PythonSettings::ProviderDefaults.snapshot(namespace);
let error =
fold_module_globals(py, Map::new(), &["vertex_project"], &snapshot).unwrap_err();
assert!(error.is_instance_of::<PyValueError>(py));
assert!(
error
.to_string()
.contains("provider_defaults.vertex_project")
);
assert!(
error
.value(py)
.getattr_opt(SETTINGS_ERROR_MARKER)
.unwrap()
.is_some()
);
});
}
#[rstest::rstest]
#[case::not_consumed(&["pages"], serde_json::json!({}), &[])]
#[case::consumed_and_absent(
&["vertex_project", "vertex_ai_project", "vertex_location", "vertex_ai_location"],
serde_json::json!({}),
&["vertex_project", "vertex_location"]
)]
#[case::consumed_and_present_in_either_spelling(
&["vertex_project", "vertex_ai_project", "vertex_location", "vertex_ai_location"],
serde_json::json!({"vertex_ai_project": "p"}),
&["vertex_location"]
)]
fn only_the_globals_of_consumed_absent_params_are_read(
#[case] consumed: &[&str],
#[case] params: Value,
#[case] expected: &[&str],
) {
assert_eq!(
module_globals_to_read(params.as_object().unwrap(), consumed),
expected
);
}
fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> {
let locals = PyDict::new(py);
py.run(source, Some(&locals), Some(&locals)).unwrap();

View file

@ -527,8 +527,6 @@ async def test_native_failures_raise_the_public_exception_class(
("ssl_verify", object()),
("ssl_certificate", 1),
("ssl_certificate", ""),
("vertex_project", 1),
("vertex_location", ["region"]),
("user_url_allowed_hosts", ["example.test", 1]),
],
)
@ -541,11 +539,35 @@ async def test_native_settings_fail_before_provider_io(
) -> None:
ocr_server.expected_requests = 0
monkeypatch.setattr(litellm, name, value)
with pytest.raises(ValueError, match=r"http_settings|provider_defaults|url_policy"):
with pytest.raises(ValueError, match=r"http_settings|url_policy"):
await call_native(ocr_server, asynchronous, num_retries=0)
assert ocr_server.requests == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
@pytest.mark.parametrize(
"name,value",
[
("vertex_project", 1),
("vertex_location", ["region"]),
],
)
async def test_native_vertex_globals_fail_before_provider_io_only_for_vertex_calls(
ocr_server: RecordingServer,
monkeypatch: pytest.MonkeyPatch,
asynchronous: bool,
name: str,
value: object,
) -> None:
ocr_server.expected_requests = 1
monkeypatch.setattr(litellm, name, value)
await call_native(ocr_server, asynchronous, num_retries=0)
with pytest.raises(ValueError, match=r"provider_defaults"):
await call_native(ocr_server, asynchronous, model="vertex_ai/mistral-ocr-latest", num_retries=0)
assert len(ocr_server.requests) == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_ssl_context_is_terminal_configuration(ocr_server: RecordingServer, asynchronous: bool) -> None: