use std::time::Instant;
use secrecy::SecretString;
use tokio::sync::RwLock;
use super::{Credential, CredentialProvider};
use crate::client::BoxFuture;
use crate::error::LiterLlmError;
const DEFAULT_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform";
const METADATA_TOKEN_URL: &str = "http://169.254.169.254/computeMetadata/v1/instance/service-accounts/default/token";
const METADATA_FLAVOR_HEADER: &str = "Metadata-Flavor";
const METADATA_FLAVOR_VALUE: &str = "Google";
const EXPIRY_BUFFER_SECS: u64 = 300;
const GCP_TOKEN_LIFETIME_SECS: u64 = 3600;
struct CachedToken {
token: SecretString,
acquired_at: Instant,
expires_in_secs: u64,
}
impl CachedToken {
fn is_valid(&self) -> bool {
let elapsed = self.acquired_at.elapsed().as_secs();
elapsed + EXPIRY_BUFFER_SECS < self.expires_in_secs
}
}
pub struct VertexAdcCredentialProvider {
scope: String,
metadata_token_url: String,
use_gcp_auth_fallback: bool,
cached: RwLock<Option<CachedToken>>,
http_client: reqwest::Client,
}
impl std::fmt::Debug for VertexAdcCredentialProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("VertexAdcCredentialProvider")
.field("scope", &self.scope)
.finish_non_exhaustive()
}
}
impl VertexAdcCredentialProvider {
#[must_use]
pub fn new() -> Self {
let scope = std::env::var("VERTEX_AI_SCOPE").unwrap_or_else(|_| DEFAULT_SCOPE.to_owned());
Self {
scope,
metadata_token_url: METADATA_TOKEN_URL.to_owned(),
use_gcp_auth_fallback: true,
cached: RwLock::new(None),
http_client: reqwest::Client::new(),
}
}
#[must_use]
pub fn with_scope(mut self, scope: impl Into<String>) -> Self {
self.scope = scope.into();
self
}
#[must_use]
pub fn with_http_client(mut self, client: reqwest::Client) -> Self {
self.http_client = client;
self
}
#[must_use]
pub fn with_metadata_url(metadata_base_url: impl Into<String>) -> Self {
let scope = std::env::var("VERTEX_AI_SCOPE").unwrap_or_else(|_| DEFAULT_SCOPE.to_owned());
let base = metadata_base_url.into();
let metadata_token_url = format!(
"{}/computeMetadata/v1/instance/service-accounts/default/token",
base.trim_end_matches('/')
);
Self {
scope,
metadata_token_url,
use_gcp_auth_fallback: true,
cached: RwLock::new(None),
http_client: reqwest::Client::new(),
}
}
#[must_use]
pub fn without_gcp_auth_fallback(mut self) -> Self {
self.use_gcp_auth_fallback = false;
self
}
async fn fetch_from_metadata_server(&self) -> Option<CachedToken> {
let response = self
.http_client
.get(&self.metadata_token_url)
.header(METADATA_FLAVOR_HEADER, METADATA_FLAVOR_VALUE)
.send()
.await
.ok()?;
if !response.status().is_success() {
#[cfg(feature = "tracing")]
tracing::warn!(
status = response.status().as_u16(),
"metadata server returned non-success status; will try gcp_auth ADC fallback"
);
return None;
}
let body = response.text().await.ok()?;
let parsed: MetadataTokenResponse = serde_json::from_str(&body).ok()?;
#[cfg(feature = "tracing")]
tracing::info!("obtained access token from metadata server");
Some(CachedToken {
token: SecretString::from(parsed.access_token),
acquired_at: Instant::now(),
expires_in_secs: parsed.expires_in,
})
}
async fn fetch_from_gcp_auth(&self) -> Result<CachedToken, LiterLlmError> {
let provider = gcp_auth::provider().await.map_err(|e| LiterLlmError::Authentication {
message: format!("gcp_auth ADC discovery failed: {e}"),
status: 401,
})?;
let scopes = &[self.scope.as_str()];
let token = provider
.token(scopes)
.await
.map_err(|e| LiterLlmError::Authentication {
message: format!("gcp_auth token acquisition failed: {e}"),
status: 401,
})?;
#[cfg(feature = "tracing")]
tracing::info!("obtained access token via gcp_auth ADC discovery");
Ok(CachedToken {
token: SecretString::from(token.as_str().to_owned()),
acquired_at: Instant::now(),
expires_in_secs: GCP_TOKEN_LIFETIME_SECS,
})
}
async fn fetch_token(&self) -> Result<CachedToken, LiterLlmError> {
if let Some(cached) = self.fetch_from_metadata_server().await {
return Ok(cached);
}
if self.use_gcp_auth_fallback {
#[cfg(feature = "tracing")]
tracing::debug!("metadata server not available; trying gcp_auth ADC discovery");
self.fetch_from_gcp_auth().await
} else {
Err(LiterLlmError::Authentication {
message: "Vertex AI ADC: metadata server unavailable and gcp_auth fallback is disabled".into(),
status: 401,
})
}
}
}
impl Default for VertexAdcCredentialProvider {
fn default() -> Self {
Self::new()
}
}
impl CredentialProvider for VertexAdcCredentialProvider {
fn resolve(&self) -> BoxFuture<'_, crate::error::Result<Credential>> {
Box::pin(async move {
{
let guard = self.cached.read().await;
if let Some(ref cached) = *guard
&& cached.is_valid()
{
#[cfg(feature = "tracing")]
tracing::debug!("returning cached Vertex AI ADC token");
return Ok(Credential::BearerToken(cached.token.clone()));
}
}
let mut guard = self.cached.write().await;
if let Some(ref cached) = *guard
&& cached.is_valid()
{
#[cfg(feature = "tracing")]
tracing::debug!("returning cached Vertex AI ADC token (post-lock check)");
return Ok(Credential::BearerToken(cached.token.clone()));
}
let fresh = self.fetch_token().await?;
let token = fresh.token.clone();
*guard = Some(fresh);
Ok(Credential::BearerToken(token))
})
}
}
#[derive(serde::Deserialize)]
struct MetadataTokenResponse {
access_token: String,
expires_in: u64,
}
#[cfg(test)]
mod tests {
use std::time::Instant;
use secrecy::SecretString;
use super::*;
#[test]
fn cached_token_is_valid_with_plenty_of_time() {
let cached = CachedToken {
token: SecretString::from("tok".to_owned()),
acquired_at: Instant::now(),
expires_in_secs: 3600,
};
assert!(cached.is_valid());
}
#[test]
fn cached_token_is_expired_at_zero_lifetime() {
let cached = CachedToken {
token: SecretString::from("tok".to_owned()),
acquired_at: Instant::now(),
expires_in_secs: 0,
};
assert!(!cached.is_valid());
}
#[test]
fn cached_token_is_expired_within_buffer() {
let cached = CachedToken {
token: SecretString::from("tok".to_owned()),
acquired_at: Instant::now(),
expires_in_secs: 200,
};
assert!(!cached.is_valid());
}
#[test]
#[serial_test::serial(vertex_adc_env)]
fn default_scope_is_cloud_platform() {
let _guard = EnvGuard::new("VERTEX_AI_SCOPE", None);
let provider = VertexAdcCredentialProvider::new();
assert_eq!(provider.scope, DEFAULT_SCOPE);
}
#[test]
#[serial_test::serial(vertex_adc_env)]
fn scope_override_via_env_var() {
let _guard = EnvGuard::new("VERTEX_AI_SCOPE", Some("https://custom.scope/"));
let provider = VertexAdcCredentialProvider::new();
assert_eq!(provider.scope, "https://custom.scope/");
}
#[test]
#[serial_test::serial(vertex_adc_env)]
fn with_scope_overrides_scope() {
let _guard = EnvGuard::new("VERTEX_AI_SCOPE", None);
let provider = VertexAdcCredentialProvider::new().with_scope("https://my.scope/");
assert_eq!(provider.scope, "https://my.scope/");
}
#[test]
#[serial_test::serial(vertex_adc_env)]
fn default_impl_equals_new() {
let _guard = EnvGuard::new("VERTEX_AI_SCOPE", None);
let provider: VertexAdcCredentialProvider = Default::default();
assert_eq!(provider.scope, DEFAULT_SCOPE);
}
#[test]
#[serial_test::serial(vertex_adc_env)]
fn with_metadata_url_appends_token_path() {
let _guard = EnvGuard::new("VERTEX_AI_SCOPE", None);
let provider = VertexAdcCredentialProvider::with_metadata_url("http://127.0.0.1:12345");
assert_eq!(
provider.metadata_token_url,
"http://127.0.0.1:12345/computeMetadata/v1/instance/service-accounts/default/token"
);
}
#[test]
#[serial_test::serial(vertex_adc_env)]
fn with_metadata_url_trailing_slash_is_normalised() {
let _guard = EnvGuard::new("VERTEX_AI_SCOPE", None);
let provider = VertexAdcCredentialProvider::with_metadata_url("http://127.0.0.1:12345/");
assert_eq!(
provider.metadata_token_url,
"http://127.0.0.1:12345/computeMetadata/v1/instance/service-accounts/default/token"
);
}
struct EnvGuard {
key: &'static str,
original: Option<String>,
}
impl EnvGuard {
fn new(key: &'static str, value: Option<&str>) -> Self {
let original = std::env::var(key).ok();
unsafe {
match value {
Some(v) => std::env::set_var(key, v),
None => std::env::remove_var(key),
}
}
Self { key, original }
}
}
impl Drop for EnvGuard {
fn drop(&mut self) {
unsafe {
match &self.original {
Some(v) => std::env::set_var(self.key, v),
None => std::env::remove_var(self.key),
}
}
}
}
#[tokio::test]
#[ignore]
async fn live_metadata_server_or_adc_returns_bearer_token() {
let provider = VertexAdcCredentialProvider::new();
let credential = provider.resolve().await.expect("token acquisition failed");
assert!(matches!(credential, Credential::BearerToken(_)));
}
}