diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index c1ee2c2e2f3..4cbde047a71 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3411,9 +3411,12 @@ dependencies = [ "litellm-auth-types", "moka", "rstest", + "serde", "serde_json", "sha2 0.10.9", + "tempfile", "tokio", + "veil", ] [[package]] diff --git a/litellm-rust/crates/auth-aws/src/constants.rs b/litellm-rust/crates/auth-aws/src/constants.rs index 61a379f8d66..789fde156bd 100644 --- a/litellm-rust/crates/auth-aws/src/constants.rs +++ b/litellm-rust/crates/auth-aws/src/constants.rs @@ -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. diff --git a/litellm-rust/crates/auth-gcp/Cargo.toml b/litellm-rust/crates/auth-gcp/Cargo.toml index 8a3598234e1..24791ce7b53 100644 --- a/litellm-rust/crates/auth-gcp/Cargo.toml +++ b/litellm-rust/crates/auth-gcp/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/auth-gcp/src/auth.rs b/litellm-rust/crates/auth-gcp/src/auth.rs new file mode 100644 index 00000000000..d7c434fe26e --- /dev/null +++ b/litellm-rust/crates/auth-gcp/src/auth.rs @@ -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>, + loader: Arc, +} + +impl Default for VertexAuth { + fn default() -> Self { + Self::new(Arc::new(GcpProviderLoader)) + } +} + +impl VertexAuth { + pub fn new(loader: Arc) -> Self { + Self { + providers: Cache::builder().max_capacity(64).build(), + loader, + } + } + + pub async fn access_token( + &self, + config: &VertexConfig, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + ) -> Result { + 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 + Sync), + ) -> Result { + 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 + Sync), + ) -> Result { + 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 + Sync), + ) -> Result, 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>; +} + +pub type VertexAuthFuture<'a, T> = Pin> + Send + 'a>>; + +struct GcpTokenSource(Arc); + +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> { + Box::pin(async move { + let provider: Arc = 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) + }) + } +} + +fn validate_request_credentials(configured: &str) -> Result<&str, Error> { + let token_uri = serde_json::from_str::(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() + ); + } +} diff --git a/litellm-rust/crates/auth-gcp/src/config.rs b/litellm-rust/crates/auth-gcp/src/config.rs new file mode 100644 index 00000000000..0a76bf93cf0 --- /dev/null +++ b/litellm-rust/crates/auth-gcp/src/config.rs @@ -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>, + project_id: Option, + location: Option, +} + +impl VertexConfig { + pub fn new( + credentials: Option>, + project_id: Option, + location: Option, + ) -> 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, + sources: &BTreeMap, + ) -> Result { + 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, +) -> Option { + 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, +) -> Option { + 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::(&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, +) -> Option { + 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, +) -> 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, + sources: &BTreeMap, + names: &[&str], +) -> Result>, 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, name: &str) -> InputSource { + sources.get(name).copied().unwrap_or_default() +} + +fn optional_string(params: &Map, names: &[&str]) -> Result, 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, + name: &str, +) -> Option { + 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(¶ms); + 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 + )); + } +} diff --git a/litellm-rust/crates/auth-gcp/src/constants.rs b/litellm-rust/crates/auth-gcp/src/constants.rs new file mode 100644 index 00000000000..e76b8dfa8ff --- /dev/null +++ b/litellm-rust/crates/auth-gcp/src/constants.rs @@ -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> = 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 +} diff --git a/litellm-rust/crates/auth-gcp/src/lib.rs b/litellm-rust/crates/auth-gcp/src/lib.rs index 97bc2c482c3..2b389d45926 100644 --- a/litellm-rust/crates/auth-gcp/src/lib.rs +++ b/litellm-rust/crates/auth-gcp/src/lib.rs @@ -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>, - project_id: Option, - location: Option, -} - -impl VertexConfig { - pub fn new( - credentials: Option>, - project_id: Option, - location: Option, - ) -> 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, - sources: &BTreeMap, - ) -> Result { - 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, -) -> Option { - 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, -) -> Option { - 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>, - loader: Arc, -} - -impl Default for VertexAuth { - fn default() -> Self { - Self::new(Arc::new(GcpProviderLoader)) - } -} - -impl VertexAuth { - pub fn new(loader: Arc) -> Self { - Self { - providers: Cache::builder().max_capacity(64).build(), - loader, - } - } - - pub async fn access_token( - &self, - config: &VertexConfig, - env_lookup: &(dyn Fn(&str) -> Option + Sync), - ) -> Result { - 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 + Sync), - ) -> Result { - 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 + Sync), - ) -> Result { - 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 + Sync), - ) -> Result, 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>; -} - -pub type VertexAuthFuture<'a, T> = Pin> + Send + 'a>>; - -struct GcpTokenSource(Arc); - -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> { - Box::pin(async move { - let provider: Arc = 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) - }) - } -} - -fn validate_request_credentials(configured: &str) -> Result<&str, Error> { - let token_uri = serde_json::from_str::(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, -) -> 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, - sources: &BTreeMap, - names: &[&str], -) -> Result>, 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, name: &str) -> InputSource { - sources.get(name).copied().unwrap_or_default() -} - -fn optional_string(params: &Map, names: &[&str]) -> Result, 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, name: &str) -> Option { - 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, - } - - 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, - provider: Arc, - } - - impl VertexProviderLoader for FakeLoader { - fn load( - &self, - _source: CredentialSource, - ) -> VertexAuthFuture<'_, Arc> { - 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, loads: Arc) -> VertexAuth { - let provider: Arc = 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::::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") - ); - } -} diff --git a/litellm-rust/crates/auth-gcp/tests/vertex_auth.rs b/litellm-rust/crates/auth-gcp/tests/vertex_auth.rs new file mode 100644 index 00000000000..5991749ab8a --- /dev/null +++ b/litellm-rust/crates/auth-gcp/tests/vertex_auth.rs @@ -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, +} + +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, + provider: Arc, +} + +impl VertexProviderLoader for FakeLoader { + fn load(&self, _source: CredentialSource) -> VertexAuthFuture<'_, Arc> { + 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, loads: Arc) -> VertexAuth { + let provider: Arc = 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::::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); +} diff --git a/litellm-rust/crates/auth-gcp/tests/vertex_config.rs b/litellm-rust/crates/auth-gcp/tests/vertex_config.rs new file mode 100644 index 00000000000..3a589ab5266 --- /dev/null +++ b/litellm-rust/crates/auth-gcp/tests/vertex_config.rs @@ -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 + ); +} diff --git a/litellm-rust/crates/auth-types/src/lib.rs b/litellm-rust/crates/auth-types/src/lib.rs index 38b13988448..c875bda7979 100644 --- a/litellm-rust/crates/auth-types/src/lib.rs +++ b/litellm-rust/crates/auth-types/src/lib.rs @@ -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}; diff --git a/litellm-rust/crates/auth-types/src/params/mod.rs b/litellm-rust/crates/auth-types/src/params/mod.rs index 093cdb42d64..88bc64cc2fc 100644 --- a/litellm-rust/crates/auth-types/src/params/mod.rs +++ b/litellm-rust/crates/auth-types/src/params/mod.rs @@ -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.` module diff --git a/litellm-rust/crates/auth-types/src/params/vertex.rs b/litellm-rust/crates/auth-types/src/params/vertex.rs new file mode 100644 index 00000000000..06ec3008dfb --- /dev/null +++ b/litellm-rust/crates/auth-types/src/params/vertex.rs @@ -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, + #[serde(default)] + pub vertex_project: Option, + #[serde(default)] + pub vertex_location: Option, + #[serde(default, deserialize_with = "optional_credential_text")] + #[redact(with = "[REDACTED]")] + pub vertex_ai_credentials: Option, + #[serde(default)] + pub vertex_ai_project: Option, + #[serde(default)] + pub vertex_ai_location: Option, +} + +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 { + 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 { + 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, + ) -> Option { + spec.resolve(&|name| self.get(name), env_lookup) + } + + pub fn credentials(&self) -> Option { + self.resolve(&Self::CREDENTIALS, &|_| None) + } + + pub fn project(&self) -> Option { + self.resolve(&Self::PROJECT, &|_| None) + } + + pub fn location(&self) -> Option { + 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, D::Error> { + match Option::::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", + } +} diff --git a/litellm-rust/crates/auth-types/tests/vertex_params.rs b/litellm-rust/crates/auth-types/tests/vertex_params.rs new file mode 100644 index 00000000000..cf0b7f824d1 --- /dev/null +++ b/litellm-rust/crates/auth-types/tests/vertex_params.rs @@ -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::(params).is_err()); +} + +#[rstest] +fn fields_name_every_param_once() { + let filled: Value = VertexParams::fields() + .map(|name| (name.to_string(), Value::from(name))) + .collect::>() + .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::>().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 + ); +} diff --git a/litellm-rust/crates/inference-ocr/tests/ocr/vertex_ai.rs b/litellm-rust/crates/inference-ocr/tests/ocr/vertex_ai.rs index 47ef3847d1f..81b52d076a5 100644 --- a/litellm-rust/crates/inference-ocr/tests/ocr/vertex_ai.rs +++ b/litellm-rust/crates/inference-ocr/tests/ocr/vertex_ai.rs @@ -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; diff --git a/litellm-rust/crates/inference-ocr/tests/resources.rs b/litellm-rust/crates/inference-ocr/tests/resources.rs index 8a5090bf54c..d618647e89b 100644 --- a/litellm-rust/crates/inference-ocr/tests/resources.rs +++ b/litellm-rust/crates/inference-ocr/tests/resources.rs @@ -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(); diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/settings.rs b/litellm-rust/crates/llms/src/base_llm/ocr/settings.rs index 87461cd36aa..af4037ddc0b 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/settings.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/settings.rs @@ -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, - pub vertex_location: Option, 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, } } diff --git a/litellm-rust/crates/llms/src/vertex_ai/ocr/common_utils.rs b/litellm-rust/crates/llms/src/vertex_ai/ocr/common_utils.rs index 93f8a25605e..810462c57f8 100644 --- a/litellm-rust/crates/llms/src/vertex_ai/ocr/common_utils.rs +++ b/litellm-rust/crates/llms/src/vertex_ai/ocr/common_utils.rs @@ -7,15 +7,10 @@ use crate::base_llm::ocr::{ }; pub(super) fn vertex_config(request: &PreparedOcrRequest) -> Result { - 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> { diff --git a/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs index 4656b2534b6..bfd7b530ffa 100644 --- a/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs @@ -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( diff --git a/litellm-rust/crates/python-bridge/src/python_settings.rs b/litellm-rust/crates/python-bridge/src/python_settings.rs index 99afc772b99..f0fea5adb06 100644 --- a/litellm-rust/crates/python-bridge/src/python_settings.rs +++ b/litellm-rust/crates/python-bridge/src/python_settings.rs @@ -68,7 +68,7 @@ fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult>()? - .is_some_and(|name| name == expected)) + .is_some_and(|name| name == expected || name.starts_with(&format!("{expected}.")))) } #[cfg(test)] diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index ca683600d71..5ae72acd104 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -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::(py) { + if !error.is_instance_of::(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; diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index 854cdb25f7b..fab9ecc13c1 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -19,10 +19,6 @@ use crate::{ python_settings::{PythonSettings, Snapshot}, }; -const VERTEX_PROJECT: FieldSpec> = - FieldSpec::new("vertex_project", |field| field.falsy_optional_string()); -const VERTEX_LOCATION: FieldSpec> = - FieldSpec::new("vertex_location", |field| field.falsy_optional_string()); const ENABLE_AZURE_AD_TOKEN_REFRESH: FieldSpec = FieldSpec::new("enable_azure_ad_token_refresh", |field| { Ok(field.exact_true()) @@ -60,8 +56,6 @@ fn ocr_settings(py: Python<'_>) -> PyResult { fn project_provider_defaults(snapshot: &Snapshot<'_>) -> PyResult { 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::(py)); assert!( - error - .to_string() - .contains("provider_defaults.vertex_project") + super::project_provider_defaults(&snapshot) + .unwrap() + .enable_azure_ad_token_refresh ); }); } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 7ae5a1f3ce2..93f04966f08 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -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::>(); 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.` 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, + 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, + globals: &[&'static str], + defaults: &Snapshot<'_>, +) -> PyResult> { + 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::>>()?; + 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::>(); + let params = params.as_object().cloned().unwrap(); + let globals = module_globals_to_read(¶ms, &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::(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(); diff --git a/tests/test_litellm_rust/ocr/test_requests.py b/tests/test_litellm_rust/ocr/test_requests.py index 8767b273b68..88cdd9ff5c1 100644 --- a/tests/test_litellm_rust/ocr/test_requests.py +++ b/tests/test_litellm_rust/ocr/test_requests.py @@ -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: