mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-11 11:49:50 +08:00
refactor(workspace): enforce layered crate boundaries
This commit is contained in:
@@ -0,0 +1,485 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
|
||||
use base64::Engine as _;
|
||||
use rsa::pkcs1::DecodeRsaPrivateKey;
|
||||
use rsa::pkcs1v15::SigningKey;
|
||||
use rsa::pkcs8::DecodePrivateKey;
|
||||
use rsa::signature::{SignatureEncoding, Signer};
|
||||
use rsa::RsaPrivateKey;
|
||||
use serde_json::{json, Value};
|
||||
use sha2::Sha256;
|
||||
use url::form_urlencoded;
|
||||
|
||||
use super::super::oauth_refresh::{
|
||||
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthHttpRequest, LocalOAuthRefreshAdapter,
|
||||
LocalOAuthRefreshError, LocalResolvedOAuthRequestAuth,
|
||||
};
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
pub const VERTEX_API_KEY_QUERY_PARAM: &str = "key";
|
||||
pub const VERTEX_SERVICE_ACCOUNT_AUTH_HEADER: &str = "authorization";
|
||||
pub const VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE: &str = "vertex_ai";
|
||||
pub const GOOGLE_OAUTH_TOKEN_URL: &str = "https://oauth2.googleapis.com/token";
|
||||
const GOOGLE_CLOUD_PLATFORM_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform";
|
||||
const SERVICE_ACCOUNT_REFRESH_SKEW_SECS: u64 = 120;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct VertexApiKeyQueryAuth {
|
||||
pub name: &'static str,
|
||||
pub value: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct VertexServiceAccountAuthConfig {
|
||||
pub client_email: String,
|
||||
pub private_key: String,
|
||||
pub project_id: String,
|
||||
pub token_uri: String,
|
||||
pub region: Option<String>,
|
||||
pub model_regions: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
pub fn resolve_local_vertex_api_key_query_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<VertexApiKeyQueryAuth> {
|
||||
if !super::is_vertex_api_key_transport_context(transport) {
|
||||
return None;
|
||||
}
|
||||
|
||||
if transport.key.decrypted_auth_config.is_some() {
|
||||
return None;
|
||||
}
|
||||
|
||||
if !transport
|
||||
.key
|
||||
.auth_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("api_key")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
if secret.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(VertexApiKeyQueryAuth {
|
||||
name: VERTEX_API_KEY_QUERY_PARAM,
|
||||
value: secret.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn resolve_local_vertex_service_account_auth_config(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<VertexServiceAccountAuthConfig> {
|
||||
if !super::is_vertex_service_account_transport_context(transport) {
|
||||
return None;
|
||||
}
|
||||
parse_vertex_service_account_auth_config(transport.key.decrypted_auth_config.as_deref())
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_service_account_auth_resolution(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
resolve_local_vertex_service_account_auth_config(transport).is_some()
|
||||
}
|
||||
|
||||
pub fn parse_vertex_service_account_auth_config(
|
||||
raw: Option<&str>,
|
||||
) -> Option<VertexServiceAccountAuthConfig> {
|
||||
let raw = raw.map(str::trim).filter(|value| !value.is_empty())?;
|
||||
let value: Value = serde_json::from_str(raw).ok()?;
|
||||
parse_vertex_service_account_auth_config_value(&value)
|
||||
}
|
||||
|
||||
fn parse_vertex_service_account_auth_config_value(
|
||||
value: &Value,
|
||||
) -> Option<VertexServiceAccountAuthConfig> {
|
||||
let client_email = json_string(value.get("client_email"))?;
|
||||
let private_key = json_string(value.get("private_key"))?;
|
||||
let project_id = json_string(value.get("project_id"))?;
|
||||
let token_uri =
|
||||
json_string(value.get("token_uri")).unwrap_or_else(|| GOOGLE_OAUTH_TOKEN_URL.to_string());
|
||||
let region = json_string(value.get("region"));
|
||||
let model_regions = value
|
||||
.get("model_regions")
|
||||
.and_then(Value::as_object)
|
||||
.map(|items| {
|
||||
items
|
||||
.iter()
|
||||
.filter_map(|(model, region)| {
|
||||
let model = model.trim();
|
||||
let region = region.as_str()?.trim();
|
||||
(!model.is_empty() && !region.is_empty())
|
||||
.then(|| (model.to_string(), region.to_string()))
|
||||
})
|
||||
.collect::<BTreeMap<_, _>>()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
Some(VertexServiceAccountAuthConfig {
|
||||
client_email,
|
||||
private_key,
|
||||
project_id,
|
||||
token_uri,
|
||||
region,
|
||||
model_regions,
|
||||
})
|
||||
}
|
||||
|
||||
fn json_string(value: Option<&Value>) -> Option<String> {
|
||||
value
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct VertexServiceAccountRefreshAdapter;
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthRefreshAdapter for VertexServiceAccountRefreshAdapter {
|
||||
fn provider_type(&self) -> &'static str {
|
||||
VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE
|
||||
}
|
||||
|
||||
fn supports(&self, transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
supports_local_vertex_service_account_auth_resolution(transport)
|
||||
}
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
if !entry
|
||||
.provider_type
|
||||
.eq_ignore_ascii_case(VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE)
|
||||
{
|
||||
return None;
|
||||
}
|
||||
if service_account_token_expires_soon(entry.expires_at_unix_secs) {
|
||||
return None;
|
||||
}
|
||||
let name = entry.auth_header_name.trim();
|
||||
let value = entry.auth_header_value.trim();
|
||||
if name.is_empty() || value.is_empty() {
|
||||
return None;
|
||||
}
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: name.to_ascii_lowercase(),
|
||||
value: value.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
None
|
||||
}
|
||||
|
||||
fn should_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> bool {
|
||||
supports_local_vertex_service_account_auth_resolution(transport)
|
||||
&& entry
|
||||
.and_then(|cached| self.resolve_cached(transport, cached))
|
||||
.is_none()
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
_entry: Option<&CachedOAuthEntry>,
|
||||
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
|
||||
let Some(auth_config) = resolve_local_vertex_service_account_auth_config(transport) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let now = aether_oauth::core::current_unix_secs();
|
||||
let assertion = build_vertex_service_account_assertion(&auth_config, now)?;
|
||||
let body = form_urlencoded::Serializer::new(String::new())
|
||||
.append_pair("grant_type", "urn:ietf:params:oauth:grant-type:jwt-bearer")
|
||||
.append_pair("assertion", &assertion)
|
||||
.finish();
|
||||
let response = executor
|
||||
.execute(
|
||||
VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
transport,
|
||||
&LocalOAuthHttpRequest {
|
||||
request_id: "vertex_ai:service-account-token",
|
||||
method: reqwest::Method::POST,
|
||||
url: auth_config.token_uri.clone(),
|
||||
headers: BTreeMap::from([(
|
||||
"content-type".to_string(),
|
||||
"application/x-www-form-urlencoded".to_string(),
|
||||
)]),
|
||||
json_body: None,
|
||||
body_bytes: Some(body.into_bytes()),
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if response.status_code != 200 {
|
||||
return Err(LocalOAuthRefreshError::HttpStatus {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
status_code: response.status_code,
|
||||
body_excerpt: body_excerpt(&response.body_text),
|
||||
});
|
||||
}
|
||||
let body_json: Value = serde_json::from_str(&response.body_text).map_err(|err| {
|
||||
LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
message: format!("vertex service account token response is not JSON: {err}"),
|
||||
}
|
||||
})?;
|
||||
let access_token = json_string(body_json.get("access_token")).ok_or_else(|| {
|
||||
LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
message: "vertex service account token response missing access_token".to_string(),
|
||||
}
|
||||
})?;
|
||||
let expires_in = body_json
|
||||
.get("expires_in")
|
||||
.and_then(Value::as_u64)
|
||||
.unwrap_or(3600);
|
||||
|
||||
Ok(Some(CachedOAuthEntry {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE.to_string(),
|
||||
auth_header_name: VERTEX_SERVICE_ACCOUNT_AUTH_HEADER.to_string(),
|
||||
auth_header_value: format!("Bearer {access_token}"),
|
||||
expires_at_unix_secs: Some(now.saturating_add(expires_in)),
|
||||
metadata: Some(json!({
|
||||
"project_id": auth_config.project_id,
|
||||
"client_email": auth_config.client_email,
|
||||
})),
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_vertex_service_account_assertion(
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<String, LocalOAuthRefreshError> {
|
||||
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"RS256","typ":"JWT"}"#);
|
||||
let payload = URL_SAFE_NO_PAD.encode(
|
||||
serde_json::to_string(&json!({
|
||||
"iss": auth_config.client_email,
|
||||
"sub": auth_config.client_email,
|
||||
"scope": GOOGLE_CLOUD_PLATFORM_SCOPE,
|
||||
"aud": auth_config.token_uri,
|
||||
"iat": now_unix_secs,
|
||||
"exp": now_unix_secs.saturating_add(3600),
|
||||
}))
|
||||
.map_err(|err| LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
message: format!("vertex service account jwt payload encode failed: {err}"),
|
||||
})?,
|
||||
);
|
||||
let message = format!("{header}.{payload}");
|
||||
let private_key = decode_vertex_service_account_private_key(auth_config.private_key.as_str())?;
|
||||
let signing_key = SigningKey::<Sha256>::new(private_key);
|
||||
let signature = signing_key.sign(message.as_bytes());
|
||||
Ok(format!(
|
||||
"{message}.{}",
|
||||
URL_SAFE_NO_PAD.encode(signature.to_bytes())
|
||||
))
|
||||
}
|
||||
|
||||
fn decode_vertex_service_account_private_key(
|
||||
private_key_pem: &str,
|
||||
) -> Result<RsaPrivateKey, LocalOAuthRefreshError> {
|
||||
match RsaPrivateKey::from_pkcs8_pem(private_key_pem) {
|
||||
Ok(private_key) => Ok(private_key),
|
||||
Err(pkcs8_err) => RsaPrivateKey::from_pkcs1_pem(private_key_pem).map_err(|pkcs1_err| {
|
||||
LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
message: format!(
|
||||
"vertex service account private_key parse failed: pkcs8: {pkcs8_err}; pkcs1: {pkcs1_err}"
|
||||
),
|
||||
}
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn service_account_token_expires_soon(expires_at_unix_secs: Option<u64>) -> bool {
|
||||
expires_at_unix_secs
|
||||
.map(|expires_at_unix_secs| {
|
||||
aether_oauth::core::current_unix_secs()
|
||||
>= expires_at_unix_secs.saturating_sub(SERVICE_ACCOUNT_REFRESH_SKEW_SECS)
|
||||
})
|
||||
.unwrap_or(true)
|
||||
}
|
||||
|
||||
fn body_excerpt(value: &str) -> String {
|
||||
value.chars().take(500).collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use rsa::pkcs1::{EncodeRsaPrivateKey, LineEnding};
|
||||
use rsa::rand_core::OsRng;
|
||||
use rsa::RsaPrivateKey;
|
||||
|
||||
use super::{
|
||||
decode_vertex_service_account_private_key, parse_vertex_service_account_auth_config,
|
||||
resolve_local_vertex_api_key_query_auth,
|
||||
supports_local_vertex_service_account_auth_resolution, VERTEX_API_KEY_QUERY_PARAM,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Vertex".to_string(),
|
||||
provider_type: "vertex_ai".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "gemini:generate_content".to_string(),
|
||||
api_family: Some("gemini".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://aiplatform.googleapis.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "api_key".to_string(),
|
||||
is_active: true,
|
||||
api_formats: Some(vec!["gemini:generate_content".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "vertex-secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_query_auth_for_vertex_api_key_subset() {
|
||||
let auth = resolve_local_vertex_api_key_query_auth(&sample_transport())
|
||||
.expect("vertex api key query auth should resolve");
|
||||
assert_eq!(auth.name, VERTEX_API_KEY_QUERY_PARAM);
|
||||
assert_eq!(auth.value, "vertex-secret");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_non_api_key_transport() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
assert!(resolve_local_vertex_api_key_query_auth(&transport).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_vertex_auth_config_transport() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_auth_config = Some("{\"project_id\":\"demo-project\"}".to_string());
|
||||
assert!(resolve_local_vertex_api_key_query_auth(&transport).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_query_auth_for_custom_aiplatform_transport() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "custom".to_string();
|
||||
transport.endpoint.api_format = "gemini:generate_content".to_string();
|
||||
|
||||
let auth = resolve_local_vertex_api_key_query_auth(&transport)
|
||||
.expect("custom aiplatform transport should resolve");
|
||||
assert_eq!(auth.value, "vertex-secret");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_service_account_auth_config() {
|
||||
let config = parse_vertex_service_account_auth_config(Some(
|
||||
r#"{
|
||||
"client_email":"[email protected]",
|
||||
"private_key":"TEST-PRIVATE-KEY",
|
||||
"project_id":"demo-project",
|
||||
"region":"global",
|
||||
"model_regions":{"gemini-2.0-flash":"us-central1"}
|
||||
}"#,
|
||||
))
|
||||
.expect("service account config should parse");
|
||||
|
||||
assert_eq!(config.client_email, "[email protected]");
|
||||
assert_eq!(config.project_id, "demo-project");
|
||||
assert_eq!(config.region.as_deref(), Some("global"));
|
||||
assert_eq!(
|
||||
config
|
||||
.model_regions
|
||||
.get("gemini-2.0-flash")
|
||||
.map(String::as_str),
|
||||
Some("us-central1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_vertex_service_account_auth_resolution() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"client_email":"[email protected]",
|
||||
"private_key":"TEST-PRIVATE-KEY",
|
||||
"project_id":"demo-project"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(supports_local_vertex_service_account_auth_resolution(
|
||||
&transport
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decodes_pkcs1_service_account_private_key() {
|
||||
let mut rng = OsRng;
|
||||
let private_key = RsaPrivateKey::new(&mut rng, 1024)
|
||||
.expect("test RSA private key should generate")
|
||||
.to_pkcs1_pem(LineEnding::LF)
|
||||
.expect("test RSA private key should encode as PKCS#1 PEM");
|
||||
|
||||
decode_vertex_service_account_private_key(private_key.as_str())
|
||||
.expect("PKCS#1 private key should decode");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,239 @@
|
||||
use url::Url;
|
||||
|
||||
use super::super::auth::resolve_local_auth_type_for_transport_format;
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
const VERTEX_AI_HOST: &str = "aiplatform.googleapis.com";
|
||||
|
||||
pub fn looks_like_vertex_ai_host(base_url: &str) -> bool {
|
||||
let trimmed = base_url.trim();
|
||||
if trimmed.is_empty() {
|
||||
return false;
|
||||
}
|
||||
|
||||
let Ok(parsed) = Url::parse(trimmed) else {
|
||||
return false;
|
||||
};
|
||||
let Some(host) = parsed
|
||||
.host_str()
|
||||
.map(|value| value.trim().to_ascii_lowercase())
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
|
||||
host == VERTEX_AI_HOST
|
||||
|| host.ends_with(&format!(".{VERTEX_AI_HOST}"))
|
||||
|| host.ends_with(&format!("-{VERTEX_AI_HOST}"))
|
||||
}
|
||||
|
||||
pub fn is_vertex_api_key_transport_context(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
if is_vertex_provider_type(transport) {
|
||||
return resolve_local_auth_type_for_transport_format(transport)
|
||||
.eq_ignore_ascii_case("api_key");
|
||||
}
|
||||
|
||||
if !is_vertex_host_format_context(transport) {
|
||||
return false;
|
||||
}
|
||||
|
||||
resolve_local_auth_type_for_transport_format(transport).eq_ignore_ascii_case("api_key")
|
||||
}
|
||||
|
||||
pub fn is_vertex_service_account_transport_context(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
if !is_vertex_provider_type(transport) && !is_vertex_host_format_context(transport) {
|
||||
return false;
|
||||
}
|
||||
|
||||
matches!(
|
||||
resolve_local_auth_type_for_transport_format(transport).as_str(),
|
||||
"service_account" | "vertex_ai"
|
||||
)
|
||||
}
|
||||
|
||||
pub fn is_vertex_transport_context(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
is_vertex_api_key_transport_context(transport)
|
||||
|| is_vertex_service_account_transport_context(transport)
|
||||
}
|
||||
|
||||
pub fn uses_vertex_api_key_query_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
provider_api_format: &str,
|
||||
) -> bool {
|
||||
is_vertex_api_key_transport_context(transport)
|
||||
&& provider_api_format
|
||||
.trim()
|
||||
.to_ascii_lowercase()
|
||||
.starts_with("gemini:")
|
||||
}
|
||||
|
||||
fn is_vertex_provider_type(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(super::PROVIDER_TYPE)
|
||||
}
|
||||
|
||||
fn is_vertex_host_format_context(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
if !looks_like_vertex_ai_host(&transport.endpoint.base_url) {
|
||||
return false;
|
||||
}
|
||||
|
||||
let endpoint_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
|
||||
endpoint_api_format.starts_with("gemini:")
|
||||
|| endpoint_api_format.starts_with("claude:")
|
||||
|| (endpoint_api_format.starts_with("openai:")
|
||||
&& looks_like_vertex_openai_compat_base(&transport.endpoint.base_url))
|
||||
}
|
||||
|
||||
fn looks_like_vertex_openai_compat_base(base_url: &str) -> bool {
|
||||
let Ok(parsed) = Url::parse(base_url.trim()) else {
|
||||
return false;
|
||||
};
|
||||
parsed
|
||||
.path()
|
||||
.trim_end_matches('/')
|
||||
.ends_with("/endpoints/openapi")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
is_vertex_api_key_transport_context, is_vertex_service_account_transport_context,
|
||||
is_vertex_transport_context, looks_like_vertex_ai_host, uses_vertex_api_key_query_auth,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Vertex".to_string(),
|
||||
provider_type: "custom".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "gemini:generate_content".to_string(),
|
||||
api_family: Some("gemini".to_string()),
|
||||
endpoint_kind: Some("cli".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://aiplatform.googleapis.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: Some("/v1/publishers/google/models/{model}:{action}".to_string()),
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "api_key".to_string(),
|
||||
is_active: true,
|
||||
api_formats: Some(vec!["gemini:generate_content".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "vertex-secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detects_vertex_host() {
|
||||
assert!(looks_like_vertex_ai_host(
|
||||
"https://aiplatform.googleapis.com"
|
||||
));
|
||||
assert!(looks_like_vertex_ai_host(
|
||||
"https://us-central1-aiplatform.googleapis.com"
|
||||
));
|
||||
assert!(!looks_like_vertex_ai_host("https://example.com"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn infers_vertex_api_key_context_for_custom_aiplatform_transport() {
|
||||
assert!(is_vertex_api_key_transport_context(&sample_transport()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_non_api_key_custom_aiplatform_transport() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "bearer".to_string();
|
||||
assert!(!is_vertex_api_key_transport_context(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn infers_vertex_service_account_context_for_fixed_provider() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "vertex_ai".to_string();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
|
||||
assert!(is_vertex_service_account_transport_context(&transport));
|
||||
assert!(is_vertex_transport_context(&transport));
|
||||
assert!(!is_vertex_api_key_transport_context(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detects_vertex_query_auth_usage_for_gemini_formats() {
|
||||
let transport = sample_transport();
|
||||
assert!(uses_vertex_api_key_query_auth(
|
||||
&transport,
|
||||
"gemini:generate_content"
|
||||
));
|
||||
assert!(!uses_vertex_api_key_query_auth(
|
||||
&transport,
|
||||
"claude:messages"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn infers_vertex_service_account_context_for_openai_compat_endpoint_root() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.api_format = "openai:chat".to_string();
|
||||
transport.endpoint.base_url =
|
||||
"https://aiplatform.googleapis.com/v1/projects/project-1/locations/global/endpoints/openapi"
|
||||
.to_string();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
|
||||
assert!(is_vertex_service_account_transport_context(&transport));
|
||||
assert!(is_vertex_transport_context(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn does_not_infer_vertex_context_for_generic_openai_format_on_aiplatform_root() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.api_format = "openai:chat".to_string();
|
||||
transport.endpoint.base_url = "https://aiplatform.googleapis.com".to_string();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
|
||||
assert!(!is_vertex_service_account_transport_context(&transport));
|
||||
assert!(!is_vertex_transport_context(&transport));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
mod auth;
|
||||
mod context;
|
||||
mod policy;
|
||||
mod url;
|
||||
|
||||
pub use auth::{
|
||||
parse_vertex_service_account_auth_config, resolve_local_vertex_api_key_query_auth,
|
||||
resolve_local_vertex_service_account_auth_config,
|
||||
supports_local_vertex_service_account_auth_resolution, VertexApiKeyQueryAuth,
|
||||
VertexServiceAccountAuthConfig, VertexServiceAccountRefreshAdapter, VERTEX_API_KEY_QUERY_PARAM,
|
||||
VERTEX_SERVICE_ACCOUNT_AUTH_HEADER, VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
};
|
||||
pub use context::{
|
||||
is_vertex_api_key_transport_context, is_vertex_service_account_transport_context,
|
||||
is_vertex_transport_context, looks_like_vertex_ai_host, uses_vertex_api_key_query_auth,
|
||||
};
|
||||
pub use policy::{
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network,
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network,
|
||||
supports_local_vertex_api_key_gemini_transport,
|
||||
supports_local_vertex_api_key_gemini_transport_with_network,
|
||||
supports_local_vertex_api_key_imagen_transport,
|
||||
supports_local_vertex_api_key_imagen_transport_with_network,
|
||||
supports_local_vertex_gemini_transport_with_network,
|
||||
};
|
||||
pub use url::{
|
||||
build_vertex_api_key_gemini_content_url, build_vertex_api_key_gemini_embedding_url,
|
||||
build_vertex_api_key_imagen_content_url, build_vertex_service_account_gemini_content_url,
|
||||
build_vertex_service_account_gemini_embedding_url, resolve_vertex_service_account_region,
|
||||
VERTEX_API_KEY_BASE_URL,
|
||||
};
|
||||
|
||||
pub const PROVIDER_TYPE: &str = "vertex_ai";
|
||||
@@ -0,0 +1,356 @@
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::super::{
|
||||
body_rules_are_locally_supported, header_rules_are_locally_supported,
|
||||
resolve_transport_profile, supports_local_oauth_request_auth_resolution,
|
||||
transport_profile_is_configured, transport_proxy_is_locally_supported,
|
||||
};
|
||||
use super::auth::{
|
||||
resolve_local_vertex_api_key_query_auth, supports_local_vertex_service_account_auth_resolution,
|
||||
};
|
||||
|
||||
fn is_vertex_transport_family(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(super::PROVIDER_TYPE)
|
||||
|| super::looks_like_vertex_ai_host(&transport.endpoint.base_url)
|
||||
}
|
||||
|
||||
pub fn local_vertex_api_key_gemini_transport_unsupported_reason_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<&'static str> {
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network_impl(transport, true)
|
||||
}
|
||||
|
||||
pub fn local_vertex_gemini_transport_unsupported_reason_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<&'static str> {
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network_impl(transport, false)
|
||||
}
|
||||
|
||||
fn local_vertex_gemini_transport_unsupported_reason_with_network_impl(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
require_api_key: bool,
|
||||
) -> Option<&'static str> {
|
||||
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
|
||||
return if !transport.provider.is_active {
|
||||
Some("provider_inactive")
|
||||
} else if !transport.endpoint.is_active {
|
||||
Some("endpoint_inactive")
|
||||
} else {
|
||||
Some("key_inactive")
|
||||
};
|
||||
}
|
||||
let endpoint_api_format =
|
||||
aether_ai_formats::normalize_api_format_alias(&transport.endpoint.api_format);
|
||||
if !matches!(
|
||||
endpoint_api_format.as_str(),
|
||||
"gemini:generate_content" | "gemini:embedding"
|
||||
) {
|
||||
return Some("transport_api_format_mismatch");
|
||||
}
|
||||
if !is_vertex_transport_family(transport) {
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
|
||||
return Some("transport_header_rules_unsupported");
|
||||
}
|
||||
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
|
||||
return Some("transport_body_rules_unsupported");
|
||||
}
|
||||
let has_api_key_auth = resolve_local_vertex_api_key_query_auth(transport).is_some();
|
||||
let has_service_account_auth = supports_local_vertex_service_account_auth_resolution(transport)
|
||||
&& supports_local_oauth_request_auth_resolution(transport);
|
||||
if require_api_key {
|
||||
if !has_api_key_auth {
|
||||
return Some("transport_auth_unavailable");
|
||||
}
|
||||
} else if !has_api_key_auth && !has_service_account_auth {
|
||||
return Some("transport_auth_unavailable");
|
||||
}
|
||||
if !transport_proxy_is_locally_supported(transport) {
|
||||
return Some("transport_proxy_unsupported");
|
||||
}
|
||||
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
|
||||
{
|
||||
return Some("transport_profile_unsupported");
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_api_key_gemini_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
supports_local_vertex_api_key_same_format_transport(
|
||||
transport,
|
||||
&["gemini:generate_content"],
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_api_key_gemini_transport_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network(transport).is_none()
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_gemini_transport_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network(transport).is_none()
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_api_key_imagen_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
supports_local_vertex_api_key_same_format_transport(
|
||||
transport,
|
||||
&["gemini:generate_content"],
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_api_key_imagen_transport_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
supports_local_vertex_api_key_same_format_transport(
|
||||
transport,
|
||||
&["gemini:generate_content"],
|
||||
true,
|
||||
)
|
||||
}
|
||||
|
||||
fn supports_local_vertex_api_key_same_format_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_formats: &[&str],
|
||||
allow_network_passthrough: bool,
|
||||
) -> bool {
|
||||
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
|
||||
return false;
|
||||
}
|
||||
let endpoint_api_format =
|
||||
aether_ai_formats::normalize_api_format_alias(&transport.endpoint.api_format);
|
||||
if !api_formats
|
||||
.iter()
|
||||
.any(|api_format| endpoint_api_format.eq_ignore_ascii_case(api_format))
|
||||
{
|
||||
return false;
|
||||
}
|
||||
if !super::is_vertex_api_key_transport_context(transport) {
|
||||
return false;
|
||||
}
|
||||
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref())
|
||||
|| !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref())
|
||||
{
|
||||
return false;
|
||||
}
|
||||
if resolve_local_vertex_api_key_query_auth(transport).is_none() {
|
||||
return false;
|
||||
}
|
||||
|
||||
let has_custom_path = transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
.as_deref()
|
||||
.is_some_and(|value: &str| !value.trim().is_empty());
|
||||
if has_custom_path && !allow_network_passthrough {
|
||||
return false;
|
||||
}
|
||||
|
||||
if allow_network_passthrough {
|
||||
if !transport_proxy_is_locally_supported(transport) {
|
||||
return false;
|
||||
}
|
||||
if transport_profile_is_configured(transport)
|
||||
&& resolve_transport_profile(transport).is_none()
|
||||
{
|
||||
return false;
|
||||
}
|
||||
} else if transport.provider.proxy.is_some()
|
||||
|| transport.endpoint.proxy.is_some()
|
||||
|| transport.key.proxy.is_some()
|
||||
|| transport_profile_is_configured(transport)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
true
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
use super::{
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network,
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network,
|
||||
supports_local_vertex_api_key_gemini_transport,
|
||||
supports_local_vertex_api_key_gemini_transport_with_network,
|
||||
supports_local_vertex_gemini_transport_with_network,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Vertex".to_string(),
|
||||
provider_type: "vertex_ai".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "gemini:generate_content".to_string(),
|
||||
api_family: Some("gemini".to_string()),
|
||||
endpoint_kind: Some("generate_content".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://aiplatform.googleapis.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "api_key".to_string(),
|
||||
is_active: true,
|
||||
api_formats: Some(vec!["gemini:generate_content".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "vertex-secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_vertex_api_key_same_format_subset() {
|
||||
assert!(supports_local_vertex_api_key_gemini_transport(
|
||||
&sample_transport()
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_vertex_api_key_gemini_generate_content_subset() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.api_format = "gemini:generate_content".to_string();
|
||||
assert!(supports_local_vertex_api_key_gemini_transport(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_custom_aiplatform_gemini_generate_content_subset() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "custom".to_string();
|
||||
transport.endpoint.api_format = "gemini:generate_content".to_string();
|
||||
assert!(supports_local_vertex_api_key_gemini_transport(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_vertex_service_account_from_api_key_subset() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_auth_config = Some("{\"project_id\":\"demo-project\"}".to_string());
|
||||
assert!(!supports_local_vertex_api_key_gemini_transport(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_vertex_service_account_gemini_transport_with_network() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"client_email":"[email protected]",
|
||||
"private_key":"TEST-PRIVATE-KEY",
|
||||
"project_id":"demo-project"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(!supports_local_vertex_api_key_gemini_transport_with_network(&transport));
|
||||
assert!(supports_local_vertex_gemini_transport_with_network(
|
||||
&transport
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_vertex_service_account_gemini_embedding_transport_with_network() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.api_format = "gemini:embedding".to_string();
|
||||
transport.endpoint.endpoint_kind = Some("embedding".to_string());
|
||||
transport.key.api_formats = Some(vec!["gemini:embedding".to_string()]);
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"client_email":"[email protected]",
|
||||
"private_key":"TEST-PRIVATE-KEY",
|
||||
"project_id":"demo-project"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(supports_local_vertex_gemini_transport_with_network(
|
||||
&transport
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn allows_network_passthrough_for_custom_path_with_local_proxy_support() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.custom_path =
|
||||
Some("/v1/publishers/google/models/gemini-2.5-pro:generateContent".to_string());
|
||||
transport.key.proxy = Some(json!({"url":"http://proxy.example:8080"}));
|
||||
transport.key.fingerprint = Some(json!({"transport_profile":"chrome_136"}));
|
||||
assert!(!supports_local_vertex_api_key_gemini_transport(&transport));
|
||||
assert!(supports_local_vertex_api_key_gemini_transport_with_network(
|
||||
&transport
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reports_auth_unavailable_when_vertex_api_key_query_auth_is_missing() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
|
||||
assert_eq!(
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network(&transport),
|
||||
Some("transport_auth_unavailable")
|
||||
);
|
||||
assert_eq!(
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network(&transport),
|
||||
Some("transport_auth_unavailable")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,327 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use url::form_urlencoded;
|
||||
|
||||
use super::super::url::build_passthrough_path_url;
|
||||
use super::auth::VertexServiceAccountAuthConfig;
|
||||
|
||||
pub const VERTEX_API_KEY_BASE_URL: &str = "https://aiplatform.googleapis.com";
|
||||
|
||||
pub fn build_vertex_api_key_gemini_content_url(
|
||||
model: &str,
|
||||
stream: bool,
|
||||
api_key: &str,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let action = if stream {
|
||||
"streamGenerateContent"
|
||||
} else {
|
||||
"generateContent"
|
||||
};
|
||||
build_vertex_api_key_google_model_url(model, action, stream, api_key, request_query)
|
||||
}
|
||||
|
||||
pub fn build_vertex_api_key_imagen_content_url(
|
||||
model: &str,
|
||||
stream: bool,
|
||||
api_key: &str,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let action = if stream {
|
||||
"streamGenerateContent"
|
||||
} else {
|
||||
"generateContent"
|
||||
};
|
||||
build_vertex_api_key_google_model_url(model, action, stream, api_key, request_query)
|
||||
}
|
||||
|
||||
pub fn build_vertex_api_key_gemini_embedding_url(
|
||||
model: &str,
|
||||
api_key: &str,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
build_vertex_api_key_google_model_url(model, "predict", false, api_key, request_query)
|
||||
}
|
||||
|
||||
pub fn build_vertex_service_account_gemini_content_url(
|
||||
model: &str,
|
||||
stream: bool,
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let action = if stream {
|
||||
"streamGenerateContent"
|
||||
} else {
|
||||
"generateContent"
|
||||
};
|
||||
build_vertex_service_account_google_model_url(model, action, stream, auth_config, request_query)
|
||||
}
|
||||
|
||||
pub fn build_vertex_service_account_gemini_embedding_url(
|
||||
model: &str,
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
build_vertex_service_account_google_model_url(
|
||||
model,
|
||||
"predict",
|
||||
false,
|
||||
auth_config,
|
||||
request_query,
|
||||
)
|
||||
}
|
||||
|
||||
fn build_vertex_api_key_google_model_url(
|
||||
model: &str,
|
||||
action: &str,
|
||||
stream: bool,
|
||||
api_key: &str,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let trimmed_model = model.trim();
|
||||
let trimmed_action = action.trim();
|
||||
let trimmed_api_key = api_key.trim();
|
||||
if trimmed_model.is_empty() || trimmed_action.is_empty() || trimmed_api_key.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let path = format!("/v1/publishers/google/models/{trimmed_model}:{trimmed_action}");
|
||||
let merged_query = build_vertex_api_key_query(trimmed_api_key, request_query, stream);
|
||||
build_passthrough_path_url(VERTEX_API_KEY_BASE_URL, &path, merged_query.as_deref(), &[])
|
||||
}
|
||||
|
||||
fn build_vertex_service_account_google_model_url(
|
||||
model: &str,
|
||||
action: &str,
|
||||
stream: bool,
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let trimmed_model = model.trim();
|
||||
let trimmed_action = action.trim();
|
||||
let project_id = auth_config.project_id.trim();
|
||||
if trimmed_model.is_empty() || trimmed_action.is_empty() || project_id.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let region = resolve_vertex_service_account_region(trimmed_model, auth_config);
|
||||
let base_url = if region == "global" {
|
||||
VERTEX_API_KEY_BASE_URL.to_string()
|
||||
} else {
|
||||
format!("https://{region}-aiplatform.googleapis.com")
|
||||
};
|
||||
let path = format!(
|
||||
"/v1/projects/{project_id}/locations/{region}/publishers/google/models/{trimmed_model}:{trimmed_action}"
|
||||
);
|
||||
let merged_query = build_vertex_service_account_query(request_query, stream);
|
||||
build_passthrough_path_url(&base_url, &path, merged_query.as_deref(), &[])
|
||||
}
|
||||
|
||||
pub fn resolve_vertex_service_account_region(
|
||||
model: &str,
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
) -> String {
|
||||
let trimmed_model = model.trim();
|
||||
if let Some(region) = auth_config
|
||||
.model_regions
|
||||
.get(trimmed_model)
|
||||
.map(String::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return region.to_string();
|
||||
}
|
||||
if let Some(region) = default_vertex_model_region(trimmed_model) {
|
||||
return region.to_string();
|
||||
}
|
||||
if let Some(region) = auth_config
|
||||
.region
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return region.to_string();
|
||||
}
|
||||
"global".to_string()
|
||||
}
|
||||
|
||||
fn default_vertex_model_region(model: &str) -> Option<&'static str> {
|
||||
if model.starts_with("gemini-3.") || model == "gemini-3-pro-image-preview" {
|
||||
return Some("global");
|
||||
}
|
||||
if matches!(
|
||||
model,
|
||||
"gemini-2.0-flash"
|
||||
| "gemini-2.0-flash-exp"
|
||||
| "gemini-2.0-flash-001"
|
||||
| "gemini-2.0-pro-exp"
|
||||
| "gemini-2.0-flash-exp-image-generation"
|
||||
| "gemini-1.5-pro"
|
||||
| "gemini-1.5-pro-001"
|
||||
| "gemini-1.5-pro-002"
|
||||
| "gemini-1.5-flash"
|
||||
| "gemini-1.5-flash-001"
|
||||
| "gemini-1.5-flash-002"
|
||||
| "imagen-3.0-generate-001"
|
||||
| "imagen-3.0-fast-generate-001"
|
||||
) {
|
||||
return Some("us-central1");
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn build_vertex_api_key_query(
|
||||
api_key: &str,
|
||||
request_query: Option<&str>,
|
||||
stream: bool,
|
||||
) -> Option<String> {
|
||||
let mut merged = BTreeMap::new();
|
||||
merge_query_string(&mut merged, request_query);
|
||||
merged.remove("beta");
|
||||
merged.insert("key".to_string(), api_key.to_string());
|
||||
if stream {
|
||||
merged
|
||||
.entry("alt".to_string())
|
||||
.or_insert_with(|| "sse".to_string());
|
||||
}
|
||||
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
for (key, value) in merged {
|
||||
serializer.append_pair(&key, &value);
|
||||
}
|
||||
let query = serializer.finish();
|
||||
if query.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(query)
|
||||
}
|
||||
}
|
||||
|
||||
fn build_vertex_service_account_query(request_query: Option<&str>, stream: bool) -> Option<String> {
|
||||
let mut merged = BTreeMap::new();
|
||||
merge_query_string(&mut merged, request_query);
|
||||
merged.remove("beta");
|
||||
merged.remove("key");
|
||||
if stream {
|
||||
merged
|
||||
.entry("alt".to_string())
|
||||
.or_insert_with(|| "sse".to_string());
|
||||
}
|
||||
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
for (key, value) in merged {
|
||||
serializer.append_pair(&key, &value);
|
||||
}
|
||||
let query = serializer.finish();
|
||||
if query.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(query)
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_query_string(out: &mut BTreeMap<String, String>, query: Option<&str>) {
|
||||
let Some(query) = query.map(str::trim).filter(|value| !value.is_empty()) else {
|
||||
return;
|
||||
};
|
||||
|
||||
for (key, value) in form_urlencoded::parse(query.as_bytes()) {
|
||||
out.insert(key.into_owned(), value.into_owned());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::{
|
||||
build_vertex_api_key_gemini_content_url, build_vertex_api_key_imagen_content_url,
|
||||
build_vertex_service_account_gemini_content_url,
|
||||
};
|
||||
use crate::vertex::VertexServiceAccountAuthConfig;
|
||||
|
||||
#[test]
|
||||
fn builds_vertex_gemini_api_key_stream_url() {
|
||||
assert_eq!(
|
||||
build_vertex_api_key_gemini_content_url(
|
||||
"gemini-2.5-pro",
|
||||
true,
|
||||
"vertex-secret",
|
||||
Some("foo=bar&beta=v1")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://aiplatform.googleapis.com/v1/publishers/google/models/gemini-2.5-pro:streamGenerateContent?alt=sse&foo=bar&key=vertex-secret"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_vertex_imagen_api_key_sync_url() {
|
||||
assert_eq!(
|
||||
build_vertex_api_key_imagen_content_url(
|
||||
"imagen-3.0-generate-001",
|
||||
false,
|
||||
"vertex-secret",
|
||||
Some("view=full")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://aiplatform.googleapis.com/v1/publishers/google/models/imagen-3.0-generate-001:generateContent?key=vertex-secret&view=full"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_vertex_service_account_gemini_sync_url() {
|
||||
let auth_config = VertexServiceAccountAuthConfig {
|
||||
client_email: "[email protected]".to_string(),
|
||||
private_key: "not-used".to_string(),
|
||||
project_id: "demo-project".to_string(),
|
||||
token_uri: "https://oauth2.googleapis.com/token".to_string(),
|
||||
region: None,
|
||||
model_regions: BTreeMap::new(),
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
build_vertex_service_account_gemini_content_url(
|
||||
"gemini-3.1-pro-preview",
|
||||
false,
|
||||
&auth_config,
|
||||
Some("foo=bar&beta=1&key=client-key")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://aiplatform.googleapis.com/v1/projects/demo-project/locations/global/publishers/google/models/gemini-3.1-pro-preview:generateContent?foo=bar"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_vertex_service_account_gemini_stream_url_with_model_region_override() {
|
||||
let auth_config = VertexServiceAccountAuthConfig {
|
||||
client_email: "[email protected]".to_string(),
|
||||
private_key: "not-used".to_string(),
|
||||
project_id: "demo-project".to_string(),
|
||||
token_uri: "https://oauth2.googleapis.com/token".to_string(),
|
||||
region: Some("global".to_string()),
|
||||
model_regions: BTreeMap::from([(
|
||||
"gemini-2.0-flash".to_string(),
|
||||
"us-central1".to_string(),
|
||||
)]),
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
build_vertex_service_account_gemini_content_url(
|
||||
"gemini-2.0-flash",
|
||||
true,
|
||||
&auth_config,
|
||||
Some("foo=bar")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://us-central1-aiplatform.googleapis.com/v1/projects/demo-project/locations/us-central1/publishers/google/models/gemini-2.0-flash:streamGenerateContent?alt=sse&foo=bar"
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user