mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
0410abea8b
commit
04a95f60c1
23 changed files with 1337 additions and 788 deletions
3
litellm-rust/Cargo.lock
generated
3
litellm-rust/Cargo.lock
generated
|
|
@ -3411,9 +3411,12 @@ dependencies = [
|
|||
"litellm-auth-types",
|
||||
"moka",
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2 0.10.9",
|
||||
"tempfile",
|
||||
"tokio",
|
||||
"veil",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
274
litellm-rust/crates/auth-gcp/src/auth.rs
Normal file
274
litellm-rust/crates/auth-gcp/src/auth.rs
Normal 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()
|
||||
);
|
||||
}
|
||||
}
|
||||
311
litellm-rust/crates/auth-gcp/src/config.rs
Normal file
311
litellm-rust/crates/auth-gcp/src/config.rs
Normal 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(¶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
|
||||
));
|
||||
}
|
||||
}
|
||||
36
litellm-rust/crates/auth-gcp/src/constants.rs
Normal file
36
litellm-rust/crates/auth-gcp/src/constants.rs
Normal 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
|
||||
}
|
||||
|
|
@ -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")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
120
litellm-rust/crates/auth-gcp/tests/vertex_auth.rs
Normal file
120
litellm-rust/crates/auth-gcp/tests/vertex_auth.rs
Normal 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);
|
||||
}
|
||||
95
litellm-rust/crates/auth-gcp/tests/vertex_config.rs
Normal file
95
litellm-rust/crates/auth-gcp/tests/vertex_config.rs
Normal 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
|
||||
);
|
||||
}
|
||||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
118
litellm-rust/crates/auth-types/src/params/vertex.rs
Normal file
118
litellm-rust/crates/auth-types/src/params/vertex.rs
Normal 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",
|
||||
}
|
||||
}
|
||||
126
litellm-rust/crates/auth-types/tests/vertex_params.rs
Normal file
126
litellm-rust/crates/auth-types/tests/vertex_params.rs
Normal 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
|
||||
);
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
);
|
||||
});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(¶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::<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();
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue