mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
Merge branch 'codex/pr-416-420-integration' into aether-rust-pioneer
This commit is contained in:
2
Cargo.lock
generated
2
Cargo.lock
generated
@@ -280,9 +280,11 @@ dependencies = [
|
||||
"aether-video-tasks-core",
|
||||
"async-trait",
|
||||
"axum",
|
||||
"base64 0.22.1",
|
||||
"http",
|
||||
"regex",
|
||||
"reqwest",
|
||||
"rsa",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2",
|
||||
|
||||
@@ -6,7 +6,9 @@ use crate::ai_serving::planner::candidate_preparation::{
|
||||
use crate::ai_serving::planner::candidate_resolution::EligibleLocalExecutionCandidate;
|
||||
use crate::ai_serving::planner::spec_metadata::local_same_format_provider_spec_metadata;
|
||||
use crate::ai_serving::transport::kiro::KiroRequestAuth;
|
||||
use crate::ai_serving::transport::vertex::resolve_local_vertex_api_key_query_auth;
|
||||
use crate::ai_serving::transport::vertex::{
|
||||
is_vertex_api_key_transport_context, resolve_local_vertex_api_key_query_auth,
|
||||
};
|
||||
use crate::ai_serving::transport::SameFormatProviderRequestBehavior;
|
||||
use crate::ai_serving::{
|
||||
GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth, PlannerAppState,
|
||||
@@ -134,7 +136,10 @@ pub(super) async fn prepare_local_same_format_provider_candidate(
|
||||
return None;
|
||||
}
|
||||
};
|
||||
if behavior.is_vertex && vertex_query_auth.is_none() {
|
||||
if behavior.is_vertex
|
||||
&& is_vertex_api_key_transport_context(&transport)
|
||||
&& vertex_query_auth.is_none()
|
||||
{
|
||||
super::super::payload::mark_skipped_local_same_format_provider_candidate(
|
||||
state,
|
||||
input,
|
||||
|
||||
@@ -524,4 +524,83 @@ mod tests {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn catalog_covers_admin_route_signatures_from_route_sources() {
|
||||
let route_sources = [
|
||||
("route/admin.rs", include_str!("route/admin.rs")),
|
||||
("route/oauth.rs", include_str!("route/oauth.rs")),
|
||||
(
|
||||
"route/public_support.rs",
|
||||
include_str!("route/public_support.rs"),
|
||||
),
|
||||
(
|
||||
"route/admin/basic_families.rs",
|
||||
include_str!("route/admin/basic_families.rs"),
|
||||
),
|
||||
(
|
||||
"route/admin/endpoints_families.rs",
|
||||
include_str!("route/admin/endpoints_families.rs"),
|
||||
),
|
||||
(
|
||||
"route/admin/model_provider_families.rs",
|
||||
include_str!("route/admin/model_provider_families.rs"),
|
||||
),
|
||||
(
|
||||
"route/admin/observability_families.rs",
|
||||
include_str!("route/admin/observability_families.rs"),
|
||||
),
|
||||
(
|
||||
"route/admin/operations_families.rs",
|
||||
include_str!("route/admin/operations_families.rs"),
|
||||
),
|
||||
(
|
||||
"route/admin/provider_ops_routes.rs",
|
||||
include_str!("route/admin/provider_ops_routes.rs"),
|
||||
),
|
||||
(
|
||||
"route/admin/system_families.rs",
|
||||
include_str!("route/admin/system_families.rs"),
|
||||
),
|
||||
];
|
||||
let mut route_scopes = BTreeSet::new();
|
||||
|
||||
for (file, source) in route_sources {
|
||||
for scope in extract_admin_route_scopes(source) {
|
||||
assert!(
|
||||
is_known_management_token_permission_scope(scope),
|
||||
"missing management token permission scope {scope} referenced by {file}"
|
||||
);
|
||||
route_scopes.insert(scope);
|
||||
}
|
||||
}
|
||||
|
||||
assert!(
|
||||
!route_scopes.is_empty(),
|
||||
"admin route scope scanner did not find any route signatures"
|
||||
);
|
||||
}
|
||||
|
||||
fn extract_admin_route_scopes(source: &'static str) -> BTreeSet<&'static str> {
|
||||
let mut scopes = BTreeSet::new();
|
||||
let mut remaining = source;
|
||||
|
||||
while let Some(start) = remaining.find("\"admin:") {
|
||||
let signature_start = start + 1;
|
||||
let after_start = &remaining[signature_start..];
|
||||
let Some(end) = after_start.find('"') else {
|
||||
break;
|
||||
};
|
||||
let signature = &after_start[..end];
|
||||
let mut parts = signature.split(':');
|
||||
if parts.next() == Some("admin") {
|
||||
if let (Some(scope), None) = (parts.next(), parts.next()) {
|
||||
scopes.insert(scope);
|
||||
}
|
||||
}
|
||||
remaining = &after_start[end + 1..];
|
||||
}
|
||||
|
||||
scopes
|
||||
}
|
||||
}
|
||||
|
||||
@@ -230,7 +230,7 @@ pub(super) fn classify_oauth_route(
|
||||
"admin_proxy",
|
||||
"provider_oauth_manage",
|
||||
"batch_import_oauth",
|
||||
"admin:provider_oauth",
|
||||
"admin:pool",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
@@ -241,7 +241,7 @@ pub(super) fn classify_oauth_route(
|
||||
"admin_proxy",
|
||||
"provider_oauth_manage",
|
||||
"start_batch_import_oauth_task",
|
||||
"admin:provider_oauth",
|
||||
"admin:pool",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::GET
|
||||
@@ -252,7 +252,7 @@ pub(super) fn classify_oauth_route(
|
||||
"admin_proxy",
|
||||
"provider_oauth_manage",
|
||||
"get_batch_import_task_status",
|
||||
"admin:provider_oauth",
|
||||
"admin:pool",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use http::Uri;
|
||||
|
||||
use crate::control::management_token_required_permission;
|
||||
|
||||
use super::{classify_control_route, headers};
|
||||
|
||||
#[test]
|
||||
@@ -66,7 +68,11 @@ fn classifies_admin_provider_oauth_batch_import_task_status_as_admin_proxy_route
|
||||
);
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("admin:provider_oauth")
|
||||
Some("admin:pool")
|
||||
);
|
||||
assert_eq!(
|
||||
management_token_required_permission(&http::Method::GET, &decision).as_deref(),
|
||||
Some("admin:pool:read")
|
||||
);
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
@@ -74,51 +80,69 @@ fn classifies_admin_provider_oauth_batch_import_task_status_as_admin_proxy_route
|
||||
#[test]
|
||||
fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() {
|
||||
let headers = headers(&[]);
|
||||
for (method, path, route_kind) in [
|
||||
for (method, path, route_kind, expected_signature, expected_required_permission) in [
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/keys/key-123/complete",
|
||||
"complete_key_oauth",
|
||||
"admin:provider_oauth",
|
||||
"admin:provider_oauth:write",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/keys/key-123/refresh",
|
||||
"refresh_key_oauth",
|
||||
"admin:provider_oauth",
|
||||
"admin:provider_oauth:write",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-123/complete",
|
||||
"complete_provider_oauth",
|
||||
"admin:provider_oauth",
|
||||
"admin:provider_oauth:write",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-123/import-refresh-token",
|
||||
"import_refresh_token",
|
||||
"admin:provider_oauth",
|
||||
"admin:provider_oauth:write",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-123/batch-import",
|
||||
"batch_import_oauth",
|
||||
"admin:pool",
|
||||
"admin:pool:write",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-123/batch-import/tasks",
|
||||
"start_batch_import_oauth_task",
|
||||
"admin:pool",
|
||||
"admin:pool:write",
|
||||
),
|
||||
(
|
||||
http::Method::GET,
|
||||
"/api/admin/provider-oauth/providers/provider-123/batch-import/tasks/task-123",
|
||||
"get_batch_import_task_status",
|
||||
"admin:pool",
|
||||
"admin:pool:read",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-123/device-authorize",
|
||||
"device_authorize",
|
||||
"admin:provider_oauth",
|
||||
"admin:provider_oauth:write",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-123/device-poll",
|
||||
"device_poll",
|
||||
"admin:provider_oauth",
|
||||
"admin:provider_oauth:write",
|
||||
),
|
||||
] {
|
||||
let uri: Uri = path.parse().expect("uri should parse");
|
||||
@@ -133,7 +157,11 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() {
|
||||
assert_eq!(decision.route_kind.as_deref(), Some(route_kind));
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("admin:provider_oauth")
|
||||
Some(expected_signature)
|
||||
);
|
||||
assert_eq!(
|
||||
management_token_required_permission(&method, &decision).as_deref(),
|
||||
Some(expected_required_permission)
|
||||
);
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
@@ -362,8 +362,8 @@ fn provider_query_transport_supports_standard_test_execution(
|
||||
)
|
||||
}
|
||||
"gemini:generate_content" => {
|
||||
if crate::provider_transport::is_vertex_api_key_transport_context(transport) {
|
||||
aether_provider_transport::vertex::supports_local_vertex_api_key_gemini_transport_with_network(transport)
|
||||
if crate::provider_transport::is_vertex_transport_context(transport) {
|
||||
aether_provider_transport::vertex::supports_local_vertex_gemini_transport_with_network(transport)
|
||||
} else {
|
||||
state.supports_local_gemini_transport_with_network(transport, api_format)
|
||||
}
|
||||
|
||||
@@ -194,7 +194,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
}
|
||||
|
||||
let oauth_auth = match format_value.as_str() {
|
||||
"openai:chat" | "claude:messages" => {
|
||||
"openai:chat" | "claude:messages" | "gemini:generate_content" => {
|
||||
match state.resolve_local_oauth_request_auth(&transport).await {
|
||||
Ok(Some(crate::provider_transport::LocalResolvedOAuthRequestAuth::Header {
|
||||
name,
|
||||
@@ -217,16 +217,26 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
}
|
||||
"gemini:generate_content" => {
|
||||
crate::provider_transport::auth::resolve_local_gemini_auth(&transport)
|
||||
.or(oauth_auth.clone())
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
let Some((auth_header, auth_value)) = auth else {
|
||||
return None;
|
||||
};
|
||||
let uses_vertex_query_auth = crate::provider_transport::uses_vertex_api_key_query_auth(
|
||||
&transport,
|
||||
format_value.as_str(),
|
||||
);
|
||||
let vertex_query_auth = if uses_vertex_query_auth {
|
||||
crate::provider_transport::vertex::resolve_local_vertex_api_key_query_auth(&transport)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let (auth_header, auth_value) = match auth {
|
||||
Some((auth_header, auth_value)) => (auth_header, auth_value),
|
||||
None if uses_vertex_query_auth && vertex_query_auth.is_some() => {
|
||||
(String::new(), String::new())
|
||||
}
|
||||
None => return None,
|
||||
};
|
||||
|
||||
let upstream_url = crate::provider_transport::build_transport_request_url(
|
||||
&transport,
|
||||
|
||||
@@ -57,6 +57,17 @@ fn oauth_access_token_expired(expires_at_unix_secs: Option<u64>, now_unix_secs:
|
||||
expires_at_unix_secs.is_none_or(|expires_at| expires_at == 0 || expires_at <= now_unix_secs)
|
||||
}
|
||||
|
||||
fn local_oauth_refresh_entry_should_stay_memory_only(
|
||||
transport: &provider_transport::GatewayProviderTransportSnapshot,
|
||||
entry: &provider_transport::CachedOAuthEntry,
|
||||
) -> bool {
|
||||
entry
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(provider_transport::vertex::VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE)
|
||||
&& provider_transport::is_vertex_service_account_transport_context(transport)
|
||||
}
|
||||
|
||||
fn oauth_auth_config_refresh_token_fingerprint(auth_config: Option<&str>) -> Option<String> {
|
||||
let parsed = auth_config
|
||||
.map(str::trim)
|
||||
@@ -1094,6 +1105,17 @@ impl AppState {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
if local_oauth_refresh_entry_should_stay_memory_only(transport, entry) {
|
||||
tracing::info!(
|
||||
key_id = %key_id,
|
||||
provider_id = %transport.provider.id,
|
||||
provider_type = %transport.provider.provider_type,
|
||||
expires_at_unix_secs = ?entry.expires_at_unix_secs,
|
||||
"gateway local oauth refresh entry kept in memory only"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let Some(encryption_key) = self.data.encryption_key() else {
|
||||
return Ok(());
|
||||
};
|
||||
@@ -1703,4 +1725,71 @@ mod tests {
|
||||
Some("[OAUTH_EXPIRED] access token invalid".to_string()),
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_service_account_refresh_entry_stays_memory_only() {
|
||||
let transport = crate::provider_transport::GatewayProviderTransportSnapshot {
|
||||
provider: crate::provider_transport::snapshot::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: crate::provider_transport::snapshot::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: crate::provider_transport::snapshot::GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "Gemini".to_string(),
|
||||
auth_type: "service_account".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,
|
||||
decrypted_api_key: "__placeholder__".to_string(),
|
||||
decrypted_auth_config: Some("{\"project_id\":\"demo\"}".to_string()),
|
||||
},
|
||||
};
|
||||
let entry = crate::provider_transport::CachedOAuthEntry {
|
||||
provider_type: "vertex_ai".to_string(),
|
||||
auth_header_name: "authorization".to_string(),
|
||||
auth_header_value: "Bearer access-token".to_string(),
|
||||
expires_at_unix_secs: Some(4_102_444_800),
|
||||
metadata: None,
|
||||
};
|
||||
|
||||
assert!(super::local_oauth_refresh_entry_should_stay_memory_only(
|
||||
&transport, &entry
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -25,9 +25,9 @@ use http::{HeaderMap, HeaderValue, StatusCode};
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::{
|
||||
build_router_with_state, build_state_with_execution_runtime_override, sample_endpoint,
|
||||
sample_key, sample_management_token, sample_oauth_provider_config, sample_provider,
|
||||
sample_proxy_node, start_server, AppState,
|
||||
build_router_with_state, build_state_with_execution_runtime_override, hash_management_token,
|
||||
sample_endpoint, sample_key, sample_management_token, sample_oauth_provider_config,
|
||||
sample_provider, sample_proxy_node, start_server, AppState,
|
||||
};
|
||||
use crate::admin_api::{
|
||||
maybe_build_local_admin_provider_oauth_response, AdminAppState, AdminRequestContext,
|
||||
@@ -7201,6 +7201,125 @@ async fn gateway_creates_updates_and_regenerates_admin_management_token_locally_
|
||||
drop(upstream_url);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_allows_management_token_with_pool_write_for_provider_oauth_batch_import() {
|
||||
let raw_token = "ae-provider-oauth-batch-pool-write";
|
||||
let state = AppState::new().expect("gateway should build");
|
||||
let admin_user = state
|
||||
.create_local_auth_user_with_settings(
|
||||
Some("provider-oauth-pool@example.com".to_string()),
|
||||
true,
|
||||
"admin".to_string(),
|
||||
"hash".to_string(),
|
||||
"admin".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("admin user should be created")
|
||||
.expect("admin user should exist");
|
||||
let mut management_token = sample_management_token(
|
||||
"token-provider-oauth-batch-pool",
|
||||
&admin_user.id,
|
||||
"provider-oauth-pool",
|
||||
true,
|
||||
);
|
||||
management_token.token.allowed_ips = None;
|
||||
management_token.token.permissions = Some(json!(["admin:pool:read", "admin:pool:write"]));
|
||||
let management_token_repository =
|
||||
Arc::new(InMemoryManagementTokenRepository::seed_with_hashes(
|
||||
vec![management_token],
|
||||
vec![(
|
||||
hash_management_token(raw_token),
|
||||
"token-provider-oauth-batch-pool".to_string(),
|
||||
)],
|
||||
));
|
||||
|
||||
let gateway = build_router_with_state(state.with_data_state_for_tests(
|
||||
GatewayDataState::with_management_token_repository_for_tests(management_token_repository),
|
||||
));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/provider-oauth/providers/provider-123/batch-import"
|
||||
))
|
||||
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||
.bearer_auth(raw_token)
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
let status = response.status();
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE, "{payload}");
|
||||
assert_eq!(payload["detail"], "Admin provider OAuth data unavailable");
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_management_token_without_pool_write_for_provider_oauth_batch_import() {
|
||||
let raw_token = "ae-provider-oauth-batch-pool-denied";
|
||||
let state = AppState::new().expect("gateway should build");
|
||||
let admin_user = state
|
||||
.create_local_auth_user_with_settings(
|
||||
Some("provider-oauth-pool-denied@example.com".to_string()),
|
||||
true,
|
||||
"admin".to_string(),
|
||||
"hash".to_string(),
|
||||
"admin".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("admin user should be created")
|
||||
.expect("admin user should exist");
|
||||
let mut management_token = sample_management_token(
|
||||
"token-provider-oauth-batch-denied",
|
||||
&admin_user.id,
|
||||
"provider-oauth-denied",
|
||||
true,
|
||||
);
|
||||
management_token.token.allowed_ips = None;
|
||||
management_token.token.permissions = Some(json!(["admin:usage:read"]));
|
||||
let management_token_repository =
|
||||
Arc::new(InMemoryManagementTokenRepository::seed_with_hashes(
|
||||
vec![management_token],
|
||||
vec![(
|
||||
hash_management_token(raw_token),
|
||||
"token-provider-oauth-batch-denied".to_string(),
|
||||
)],
|
||||
));
|
||||
|
||||
let gateway = build_router_with_state(state.with_data_state_for_tests(
|
||||
GatewayDataState::with_management_token_repository_for_tests(management_token_repository),
|
||||
));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/provider-oauth/providers/provider-123/batch-import"
|
||||
))
|
||||
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||
.bearer_auth(raw_token)
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::FORBIDDEN);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["detail"], "management token permission denied");
|
||||
assert_eq!(payload["required_permission"], "admin:pool:write");
|
||||
assert_eq!(payload["route_family"], "provider_oauth_manage");
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_deletes_admin_management_token_locally_with_trusted_admin_principal() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
|
||||
@@ -351,7 +351,8 @@ CREATE TABLE IF NOT EXISTS public.management_tokens (
|
||||
token_prefix character varying(12),
|
||||
name character varying(100) NOT NULL,
|
||||
description text,
|
||||
allowed_ips json,
|
||||
allowed_ips jsonb,
|
||||
permissions jsonb,
|
||||
expires_at timestamp with time zone,
|
||||
last_used_at timestamp with time zone,
|
||||
last_used_ip character varying(45),
|
||||
@@ -359,7 +360,7 @@ CREATE TABLE IF NOT EXISTS public.management_tokens (
|
||||
is_active boolean DEFAULT true NOT NULL,
|
||||
created_at timestamp with time zone DEFAULT now() NOT NULL,
|
||||
updated_at timestamp with time zone DEFAULT now() NOT NULL,
|
||||
CONSTRAINT check_allowed_ips_not_empty CHECK (((allowed_ips IS NULL) OR ((allowed_ips)::text = 'null'::text) OR (json_array_length(allowed_ips) > 0)))
|
||||
CONSTRAINT check_allowed_ips_not_empty CHECK (CASE WHEN ((allowed_ips IS NULL) OR (allowed_ips = 'null'::jsonb)) THEN true WHEN (jsonb_typeof(allowed_ips) = 'array'::text) THEN (jsonb_array_length(allowed_ips) > 0) ELSE false END)
|
||||
);
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
ALTER TABLE public.management_tokens
|
||||
DROP CONSTRAINT IF EXISTS check_allowed_ips_not_empty;
|
||||
|
||||
ALTER TABLE public.management_tokens
|
||||
ADD COLUMN IF NOT EXISTS permissions jsonb;
|
||||
|
||||
ALTER TABLE public.management_tokens
|
||||
ALTER COLUMN allowed_ips TYPE jsonb USING allowed_ips::jsonb,
|
||||
ALTER COLUMN permissions TYPE jsonb USING permissions::jsonb;
|
||||
|
||||
ALTER TABLE public.management_tokens
|
||||
ADD CONSTRAINT check_allowed_ips_not_empty CHECK (
|
||||
CASE
|
||||
WHEN allowed_ips IS NULL OR allowed_ips = 'null'::jsonb THEN TRUE
|
||||
WHEN jsonb_typeof(allowed_ips) = 'array' THEN jsonb_array_length(allowed_ips) > 0
|
||||
ELSE FALSE
|
||||
END
|
||||
);
|
||||
@@ -352,7 +352,8 @@ CREATE TABLE IF NOT EXISTS public.management_tokens (
|
||||
token_prefix character varying(12),
|
||||
name character varying(100) NOT NULL,
|
||||
description text,
|
||||
allowed_ips json,
|
||||
allowed_ips jsonb,
|
||||
permissions jsonb,
|
||||
expires_at timestamp with time zone,
|
||||
last_used_at timestamp with time zone,
|
||||
last_used_ip character varying(45),
|
||||
@@ -360,7 +361,7 @@ CREATE TABLE IF NOT EXISTS public.management_tokens (
|
||||
is_active boolean DEFAULT true NOT NULL,
|
||||
created_at timestamp with time zone DEFAULT now() NOT NULL,
|
||||
updated_at timestamp with time zone DEFAULT now() NOT NULL,
|
||||
CONSTRAINT check_allowed_ips_not_empty CHECK (((allowed_ips IS NULL) OR ((allowed_ips)::text = 'null'::text) OR (json_array_length(allowed_ips) > 0)))
|
||||
CONSTRAINT check_allowed_ips_not_empty CHECK (CASE WHEN ((allowed_ips IS NULL) OR (allowed_ips = 'null'::jsonb)) THEN true WHEN (jsonb_typeof(allowed_ips) = 'array'::text) THEN (jsonb_array_length(allowed_ips) > 0) ELSE false END)
|
||||
);
|
||||
|
||||
|
||||
|
||||
@@ -351,7 +351,8 @@ CREATE TABLE IF NOT EXISTS public.management_tokens (
|
||||
token_prefix character varying(12),
|
||||
name character varying(100) NOT NULL,
|
||||
description text,
|
||||
allowed_ips json,
|
||||
allowed_ips jsonb,
|
||||
permissions jsonb,
|
||||
expires_at timestamp with time zone,
|
||||
last_used_at timestamp with time zone,
|
||||
last_used_ip character varying(45),
|
||||
@@ -359,7 +360,7 @@ CREATE TABLE IF NOT EXISTS public.management_tokens (
|
||||
is_active boolean DEFAULT true NOT NULL,
|
||||
created_at timestamp with time zone DEFAULT now() NOT NULL,
|
||||
updated_at timestamp with time zone DEFAULT now() NOT NULL,
|
||||
CONSTRAINT check_allowed_ips_not_empty CHECK (((allowed_ips IS NULL) OR ((allowed_ips)::text = 'null'::text) OR (json_array_length(allowed_ips) > 0)))
|
||||
CONSTRAINT check_allowed_ips_not_empty CHECK (CASE WHEN ((allowed_ips IS NULL) OR (allowed_ips = 'null'::jsonb)) THEN true WHEN (jsonb_typeof(allowed_ips) = 'array'::text) THEN (jsonb_array_length(allowed_ips) > 0) ELSE false END)
|
||||
);
|
||||
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ use tracing::info;
|
||||
// Generated by build.rs from schema/bootstrap/postgres.
|
||||
pub(crate) static EMPTY_DATABASE_SNAPSHOT_SQL: &str =
|
||||
include_str!(concat!(env!("OUT_DIR"), "/empty_database_snapshot.sql"));
|
||||
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260509120000;
|
||||
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260510000000;
|
||||
|
||||
const PUBLIC_BASE_TABLE_COUNT_SQL: &str = r#"
|
||||
SELECT COUNT(*)::BIGINT
|
||||
@@ -27,12 +27,12 @@ WHERE table_schema = 'public'
|
||||
'gemini_file_mappings',
|
||||
'global_models',
|
||||
'oauth_providers',
|
||||
'provider_api_keys',
|
||||
'proxy_nodes',
|
||||
'user_groups',
|
||||
'usage_routing_snapshots',
|
||||
'usage_settlement_snapshots'
|
||||
)
|
||||
'provider_api_keys',
|
||||
'proxy_nodes',
|
||||
'user_groups',
|
||||
'usage_routing_snapshots',
|
||||
'usage_settlement_snapshots'
|
||||
)
|
||||
"#;
|
||||
const INSERT_APPLIED_MIGRATION_SQL: &str = r#"
|
||||
INSERT INTO _sqlx_migrations (
|
||||
|
||||
@@ -296,6 +296,7 @@ fn empty_database_snapshot_covers_current_cutoff_versions() {
|
||||
20260508000000,
|
||||
20260509000000,
|
||||
20260509120000,
|
||||
20260510000000,
|
||||
]
|
||||
);
|
||||
}
|
||||
@@ -396,6 +397,44 @@ fn provider_api_keys_api_formats_remains_nullable_in_baselines() {
|
||||
.contains("pak.allow_auth_channel_mismatch_formats IS NULL"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn management_tokens_json_columns_are_normalized_to_jsonb_in_postgres_schema_paths() {
|
||||
let normalization_migration = POSTGRES_MIGRATOR
|
||||
.iter()
|
||||
.find(|migration| migration.version == 20260510000000)
|
||||
.expect("management token jsonb normalization migration should be embedded");
|
||||
assert!(normalization_migration
|
||||
.sql
|
||||
.contains("ALTER COLUMN allowed_ips TYPE jsonb USING allowed_ips::jsonb"));
|
||||
assert!(normalization_migration
|
||||
.sql
|
||||
.contains("ALTER COLUMN permissions TYPE jsonb USING permissions::jsonb"));
|
||||
assert!(normalization_migration
|
||||
.sql
|
||||
.contains("jsonb_array_length(allowed_ips) > 0"));
|
||||
|
||||
assert!(EMPTY_DATABASE_SNAPSHOT_SQL.contains("allowed_ips jsonb,"));
|
||||
assert!(EMPTY_DATABASE_SNAPSHOT_SQL.contains("permissions jsonb,"));
|
||||
assert!(EMPTY_DATABASE_SNAPSHOT_SQL.contains("jsonb_array_length(allowed_ips)"));
|
||||
|
||||
let bootstrap_schema =
|
||||
include_str!("../../../schema/bootstrap/postgres/001_types_and_tables.sql");
|
||||
assert!(bootstrap_schema.contains("allowed_ips jsonb,"));
|
||||
assert!(bootstrap_schema.contains("permissions jsonb,"));
|
||||
assert!(bootstrap_schema.contains("jsonb_array_length(allowed_ips)"));
|
||||
|
||||
let driver_schema =
|
||||
include_str!("../../../schema/drivers/postgres/baseline/001_types_and_tables.sql");
|
||||
assert!(driver_schema.contains("allowed_ips jsonb,"));
|
||||
assert!(driver_schema.contains("permissions jsonb,"));
|
||||
assert!(driver_schema.contains("jsonb_array_length(allowed_ips)"));
|
||||
|
||||
let generated_identity =
|
||||
include_str!("../../../schema/generated/postgres/baseline/001_identity.sql");
|
||||
assert!(generated_identity.contains("allowed_ips jsonb,"));
|
||||
assert!(generated_identity.contains("permissions jsonb,"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_api_keys_api_key_is_nullable() {
|
||||
let baseline_migration = POSTGRES_MIGRATOR
|
||||
@@ -1038,6 +1077,7 @@ fn pending_migrations_from_applied_skips_versions_already_applied() {
|
||||
20260508000000,
|
||||
20260509000000,
|
||||
20260509120000,
|
||||
20260510000000,
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
@@ -103,7 +103,15 @@ DELETE FROM management_tokens
|
||||
WHERE id = $1
|
||||
"#;
|
||||
|
||||
const CREATE_MANAGEMENT_TOKEN_SQL: &str = r#"
|
||||
const MANAGEMENT_TOKEN_JSON_COLUMN_TYPES_SQL: &str = r#"
|
||||
SELECT column_name, udt_name
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = 'public'
|
||||
AND table_name = 'management_tokens'
|
||||
AND column_name IN ('allowed_ips', 'permissions')
|
||||
"#;
|
||||
|
||||
const CREATE_MANAGEMENT_TOKEN_SQL_PREFIX: &str = r#"
|
||||
INSERT INTO management_tokens (
|
||||
id,
|
||||
user_id,
|
||||
@@ -123,8 +131,9 @@ VALUES (
|
||||
$4,
|
||||
$5,
|
||||
$6,
|
||||
$7,
|
||||
$8,
|
||||
"#;
|
||||
|
||||
const CREATE_MANAGEMENT_TOKEN_SQL_SUFFIX: &str = r#",
|
||||
CASE
|
||||
WHEN $9::bigint IS NULL THEN NULL
|
||||
ELSE to_timestamp($9::double precision)
|
||||
@@ -148,26 +157,23 @@ RETURNING
|
||||
EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs
|
||||
"#;
|
||||
|
||||
const UPDATE_MANAGEMENT_TOKEN_SQL: &str = r#"
|
||||
const UPDATE_MANAGEMENT_TOKEN_SQL_PREFIX: &str = r#"
|
||||
UPDATE management_tokens
|
||||
SET name = COALESCE($2, name),
|
||||
description = CASE
|
||||
WHEN $3 THEN NULL
|
||||
WHEN $4::text IS NULL THEN description
|
||||
ELSE $4
|
||||
END,
|
||||
allowed_ips = CASE
|
||||
WHEN $5 THEN NULL
|
||||
WHEN $6::json IS NULL THEN allowed_ips
|
||||
ELSE $6
|
||||
END,
|
||||
permissions = COALESCE($7::json, permissions),
|
||||
SET name = $2,
|
||||
description = $3,
|
||||
allowed_ips =
|
||||
"#;
|
||||
|
||||
const UPDATE_MANAGEMENT_TOKEN_SQL_MIDDLE: &str = r#",
|
||||
permissions =
|
||||
"#;
|
||||
|
||||
const UPDATE_MANAGEMENT_TOKEN_SQL_SUFFIX: &str = r#",
|
||||
expires_at = CASE
|
||||
WHEN $8 THEN NULL
|
||||
WHEN $9::bigint IS NULL THEN expires_at
|
||||
ELSE to_timestamp($9::double precision)
|
||||
WHEN $6::bigint IS NULL THEN NULL
|
||||
ELSE to_timestamp($6::double precision)
|
||||
END,
|
||||
is_active = COALESCE($10, is_active),
|
||||
is_active = $7,
|
||||
updated_at = NOW()
|
||||
WHERE id = $1
|
||||
RETURNING
|
||||
@@ -261,10 +267,93 @@ pub struct SqlxManagementTokenRepository {
|
||||
pool: PgPool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum JsonColumnType {
|
||||
Json,
|
||||
Jsonb,
|
||||
}
|
||||
|
||||
impl JsonColumnType {
|
||||
fn from_udt_name(value: &str) -> Option<Self> {
|
||||
match value {
|
||||
"json" => Some(Self::Json),
|
||||
"jsonb" => Some(Self::Jsonb),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn sql_type(self) -> &'static str {
|
||||
match self {
|
||||
Self::Json => "json",
|
||||
Self::Jsonb => "jsonb",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
struct ManagementTokenJsonColumnTypes {
|
||||
allowed_ips: JsonColumnType,
|
||||
permissions: JsonColumnType,
|
||||
}
|
||||
|
||||
impl SqlxManagementTokenRepository {
|
||||
pub fn new(pool: PgPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
async fn json_column_types(&self) -> Result<ManagementTokenJsonColumnTypes, DataLayerError> {
|
||||
let rows = sqlx::query(MANAGEMENT_TOKEN_JSON_COLUMN_TYPES_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let mut allowed_ips = None;
|
||||
let mut permissions = None;
|
||||
for row in rows {
|
||||
let column_name: String = row.try_get("column_name").map_postgres_err()?;
|
||||
let udt_name: String = row.try_get("udt_name").map_postgres_err()?;
|
||||
let Some(column_type) = JsonColumnType::from_udt_name(udt_name.as_str()) else {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"unsupported management_tokens.{column_name} column type: {udt_name}"
|
||||
)));
|
||||
};
|
||||
match column_name.as_str() {
|
||||
"allowed_ips" => allowed_ips = Some(column_type),
|
||||
"permissions" => permissions = Some(column_type),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
match (allowed_ips, permissions) {
|
||||
(Some(allowed_ips), Some(permissions)) => Ok(ManagementTokenJsonColumnTypes {
|
||||
allowed_ips,
|
||||
permissions,
|
||||
}),
|
||||
_ => Err(DataLayerError::UnexpectedValue(
|
||||
"management_tokens JSON column metadata missing".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn create_management_token_sql(types: ManagementTokenJsonColumnTypes) -> String {
|
||||
format!(
|
||||
"{} $7::text::{},\n $8::text::{}{}",
|
||||
CREATE_MANAGEMENT_TOKEN_SQL_PREFIX,
|
||||
types.allowed_ips.sql_type(),
|
||||
types.permissions.sql_type(),
|
||||
CREATE_MANAGEMENT_TOKEN_SQL_SUFFIX
|
||||
)
|
||||
}
|
||||
|
||||
fn update_management_token_sql(types: ManagementTokenJsonColumnTypes) -> String {
|
||||
format!(
|
||||
"{} $4::text::{}{} $5::text::{}{}",
|
||||
UPDATE_MANAGEMENT_TOKEN_SQL_PREFIX,
|
||||
types.allowed_ips.sql_type(),
|
||||
UPDATE_MANAGEMENT_TOKEN_SQL_MIDDLE,
|
||||
types.permissions.sql_type(),
|
||||
UPDATE_MANAGEMENT_TOKEN_SQL_SUFFIX
|
||||
)
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -330,15 +419,19 @@ impl ManagementTokenWriteRepository for SqlxManagementTokenRepository {
|
||||
record: &CreateManagementTokenRecord,
|
||||
) -> Result<StoredManagementToken, DataLayerError> {
|
||||
record.validate()?;
|
||||
let row = sqlx::query(CREATE_MANAGEMENT_TOKEN_SQL)
|
||||
let json_column_types = self.json_column_types().await?;
|
||||
let sql = create_management_token_sql(json_column_types);
|
||||
let allowed_ips = json_to_string(record.allowed_ips.as_ref())?;
|
||||
let permissions = json_to_string(record.permissions.as_ref())?;
|
||||
let row = sqlx::query(sql.as_str())
|
||||
.bind(&record.id)
|
||||
.bind(&record.user_id)
|
||||
.bind(&record.token_hash)
|
||||
.bind(record.token_prefix.as_deref())
|
||||
.bind(&record.name)
|
||||
.bind(record.description.as_deref())
|
||||
.bind(record.allowed_ips.as_ref())
|
||||
.bind(record.permissions.as_ref())
|
||||
.bind(allowed_ips)
|
||||
.bind(permissions)
|
||||
.bind(
|
||||
record
|
||||
.expires_at_unix_secs
|
||||
@@ -356,21 +449,56 @@ impl ManagementTokenWriteRepository for SqlxManagementTokenRepository {
|
||||
record: &UpdateManagementTokenRecord,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
record.validate()?;
|
||||
let row = sqlx::query(UPDATE_MANAGEMENT_TOKEN_SQL)
|
||||
let Some(current) = self
|
||||
.get_management_token_with_user(&record.token_id)
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let json_column_types = self.json_column_types().await?;
|
||||
let sql = update_management_token_sql(json_column_types);
|
||||
let name = record
|
||||
.name
|
||||
.as_deref()
|
||||
.unwrap_or(current.token.name.as_str());
|
||||
let description = if record.clear_description {
|
||||
None
|
||||
} else {
|
||||
record
|
||||
.description
|
||||
.as_deref()
|
||||
.or(current.token.description.as_deref())
|
||||
};
|
||||
let allowed_ips = if record.clear_allowed_ips {
|
||||
None
|
||||
} else {
|
||||
record
|
||||
.allowed_ips
|
||||
.as_ref()
|
||||
.or(current.token.allowed_ips.as_ref())
|
||||
};
|
||||
let permissions = record
|
||||
.permissions
|
||||
.as_ref()
|
||||
.or(current.token.permissions.as_ref());
|
||||
let expires_at_unix_secs = if record.clear_expires_at {
|
||||
None
|
||||
} else {
|
||||
record
|
||||
.expires_at_unix_secs
|
||||
.or(current.token.expires_at_unix_secs)
|
||||
};
|
||||
let is_active = record.is_active.unwrap_or(current.token.is_active);
|
||||
let allowed_ips = json_to_string(allowed_ips)?;
|
||||
let permissions = json_to_string(permissions)?;
|
||||
let row = sqlx::query(sql.as_str())
|
||||
.bind(&record.token_id)
|
||||
.bind(record.name.as_deref())
|
||||
.bind(record.clear_description)
|
||||
.bind(record.description.as_deref())
|
||||
.bind(record.clear_allowed_ips)
|
||||
.bind(record.allowed_ips.as_ref())
|
||||
.bind(record.permissions.as_ref())
|
||||
.bind(record.clear_expires_at)
|
||||
.bind(
|
||||
record
|
||||
.expires_at_unix_secs
|
||||
.and_then(|value| i64::try_from(value).ok()),
|
||||
)
|
||||
.bind(record.is_active)
|
||||
.bind(name)
|
||||
.bind(description)
|
||||
.bind(allowed_ips)
|
||||
.bind(permissions)
|
||||
.bind(expires_at_unix_secs.and_then(|value| i64::try_from(value).ok()))
|
||||
.bind(is_active)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_err(|err| map_management_token_write_error(err, record.name.as_deref()))?;
|
||||
@@ -434,6 +562,18 @@ fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
|
||||
value.and_then(|value| u64::try_from(value).ok())
|
||||
}
|
||||
|
||||
fn json_to_string(value: Option<&serde_json::Value>) -> Result<Option<String>, DataLayerError> {
|
||||
value
|
||||
.map(|value| {
|
||||
serde_json::to_string(value).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid management token JSON field: {err}"
|
||||
))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn map_management_token_write_error(
|
||||
err: sqlx::Error,
|
||||
requested_name: Option<&str>,
|
||||
@@ -503,7 +643,10 @@ fn map_token_with_user_row(row: &PgRow) -> Result<StoredManagementTokenWithUser,
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::SqlxManagementTokenRepository;
|
||||
use super::{
|
||||
create_management_token_sql, update_management_token_sql, JsonColumnType,
|
||||
ManagementTokenJsonColumnTypes, SqlxManagementTokenRepository,
|
||||
};
|
||||
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
|
||||
#[tokio::test]
|
||||
@@ -523,4 +666,42 @@ mod tests {
|
||||
let pool = factory.connect_lazy().expect("pool should build");
|
||||
let _repository = SqlxManagementTokenRepository::new(pool);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn repository_sql_casts_json_fields_to_detected_column_types_without_json_case() {
|
||||
let jsonb_types = ManagementTokenJsonColumnTypes {
|
||||
allowed_ips: JsonColumnType::Jsonb,
|
||||
permissions: JsonColumnType::Jsonb,
|
||||
};
|
||||
let json_types = ManagementTokenJsonColumnTypes {
|
||||
allowed_ips: JsonColumnType::Json,
|
||||
permissions: JsonColumnType::Json,
|
||||
};
|
||||
let jsonb_create_sql = create_management_token_sql(jsonb_types);
|
||||
let jsonb_update_sql = update_management_token_sql(jsonb_types);
|
||||
let json_create_sql = create_management_token_sql(json_types);
|
||||
let json_update_sql = update_management_token_sql(json_types);
|
||||
|
||||
assert!(jsonb_create_sql.contains("$7::text::jsonb"));
|
||||
assert!(jsonb_create_sql.contains("$8::text::jsonb"));
|
||||
assert!(jsonb_update_sql.contains("allowed_ips =\n $4::text::jsonb"));
|
||||
assert!(jsonb_update_sql.contains("permissions =\n $5::text::jsonb"));
|
||||
|
||||
assert!(json_create_sql.contains("$7::text::json"));
|
||||
assert!(json_create_sql.contains("$8::text::json"));
|
||||
assert!(json_update_sql.contains("allowed_ips =\n $4::text::json"));
|
||||
assert!(json_update_sql.contains("permissions =\n $5::text::json"));
|
||||
|
||||
for sql in [
|
||||
jsonb_create_sql.as_str(),
|
||||
jsonb_update_sql.as_str(),
|
||||
json_create_sql.as_str(),
|
||||
json_update_sql.as_str(),
|
||||
] {
|
||||
assert!(!sql.contains("allowed_ips = CASE"));
|
||||
assert!(!sql.contains("permissions = CASE"));
|
||||
assert!(!sql.contains("$6::json IS NULL"));
|
||||
assert!(!sql.contains("COALESCE($7::json"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8,6 +8,7 @@ use aether_provider_transport::{
|
||||
};
|
||||
use base64::engine::general_purpose::{STANDARD, 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};
|
||||
@@ -675,8 +676,7 @@ fn build_vertex_service_account_assertion(
|
||||
.map_err(|err| format!("vertex_ai(service_account): jwt payload encode failed: {err}"))?,
|
||||
);
|
||||
let message = format!("{header}.{payload}");
|
||||
let private_key = RsaPrivateKey::from_pkcs8_pem(private_key_pem)
|
||||
.map_err(|err| format!("vertex_ai(service_account): private_key parse failed: {err}"))?;
|
||||
let private_key = decode_vertex_service_account_private_key(private_key_pem)?;
|
||||
let signing_key = SigningKey::<Sha256>::new(private_key);
|
||||
let signature = signing_key.sign(message.as_bytes());
|
||||
Ok(format!(
|
||||
@@ -685,6 +685,19 @@ fn build_vertex_service_account_assertion(
|
||||
))
|
||||
}
|
||||
|
||||
fn decode_vertex_service_account_private_key(
|
||||
private_key_pem: &str,
|
||||
) -> Result<RsaPrivateKey, String> {
|
||||
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| {
|
||||
format!(
|
||||
"vertex_ai(service_account): private_key parse failed: pkcs8: {pkcs8_err}; pkcs1: {pkcs1_err}"
|
||||
)
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn execution_result_json_body(result: &ExecutionResult) -> Result<Value, String> {
|
||||
if result.status_code != 200 {
|
||||
return Err(execution_result_error_message(result));
|
||||
|
||||
@@ -15,12 +15,14 @@ aether-oauth.workspace = true
|
||||
aether-runtime-state.workspace = true
|
||||
aether-video-tasks-core.workspace = true
|
||||
async-trait.workspace = true
|
||||
base64.workspace = true
|
||||
http.workspace = true
|
||||
regex.workspace = true
|
||||
reqwest.workspace = true
|
||||
rsa = "0.9.10"
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
sha2.workspace = true
|
||||
sha2 = { workspace = true, features = ["oid"] }
|
||||
thiserror.workspace = true
|
||||
tokio.workspace = true
|
||||
tracing.workspace = true
|
||||
|
||||
@@ -15,8 +15,8 @@ use crate::policy::{
|
||||
local_standard_transport_unsupported_reason_with_network,
|
||||
};
|
||||
use crate::vertex::{
|
||||
is_vertex_api_key_transport_context,
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network,
|
||||
is_vertex_api_key_transport_context, is_vertex_transport_context,
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network,
|
||||
resolve_local_vertex_api_key_query_auth, VERTEX_API_KEY_QUERY_PARAM,
|
||||
};
|
||||
use crate::GatewayProviderTransportSnapshot;
|
||||
@@ -106,8 +106,8 @@ pub fn request_conversion_transport_unsupported_reason(
|
||||
"claude:messages" => {
|
||||
local_standard_transport_unsupported_reason_with_network(transport, "claude:messages")
|
||||
}
|
||||
"gemini:generate_content" if is_vertex_api_key_transport_context(transport) => {
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network(transport)
|
||||
"gemini:generate_content" if is_vertex_transport_context(transport) => {
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network(transport)
|
||||
}
|
||||
"gemini:generate_content" => local_gemini_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
|
||||
@@ -103,7 +103,10 @@ pub use standard::{
|
||||
StandardPlanFallbackAcceptPolicy, StandardPlanFallbackHeadersInput,
|
||||
StandardProviderRequestHeaders, StandardProviderRequestHeadersInput,
|
||||
};
|
||||
pub use vertex::{is_vertex_api_key_transport_context, uses_vertex_api_key_query_auth};
|
||||
pub use vertex::{
|
||||
is_vertex_api_key_transport_context, is_vertex_service_account_transport_context,
|
||||
is_vertex_transport_context, uses_vertex_api_key_query_auth,
|
||||
};
|
||||
pub use video::{
|
||||
build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url,
|
||||
reconstruct_local_video_task_snapshot, resolve_local_video_task_transport,
|
||||
|
||||
@@ -19,6 +19,9 @@ use super::kiro::{
|
||||
supports_local_kiro_request_auth_resolution, KiroOAuthRefreshAdapter, KiroRequestAuth,
|
||||
};
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::vertex::{
|
||||
supports_local_vertex_service_account_auth_resolution, VertexServiceAccountRefreshAdapter,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[allow(clippy::large_enum_variant)]
|
||||
@@ -339,6 +342,7 @@ impl LocalOAuthRefreshCoordinator {
|
||||
Self {
|
||||
adapters: vec![
|
||||
Arc::new(KiroOAuthRefreshAdapter::default()),
|
||||
Arc::new(VertexServiceAccountRefreshAdapter),
|
||||
Arc::new(GenericOAuthRefreshAdapter::default()),
|
||||
],
|
||||
cache: Mutex::new(BTreeMap::new()),
|
||||
@@ -547,6 +551,7 @@ pub fn supports_local_oauth_request_auth_resolution(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
supports_local_kiro_request_auth_resolution(transport)
|
||||
|| supports_local_vertex_service_account_auth_resolution(transport)
|
||||
|| supports_local_generic_oauth_request_auth_resolution(transport)
|
||||
}
|
||||
|
||||
|
||||
@@ -15,7 +15,8 @@ use crate::url::{
|
||||
build_openai_responses_url, build_passthrough_path_url, normalize_gemini_content_action_path,
|
||||
};
|
||||
use crate::vertex::{
|
||||
build_vertex_api_key_gemini_content_url, resolve_local_vertex_api_key_query_auth,
|
||||
build_vertex_api_key_gemini_content_url, build_vertex_service_account_gemini_content_url,
|
||||
resolve_local_vertex_api_key_query_auth, resolve_local_vertex_service_account_auth_config,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
@@ -256,6 +257,14 @@ fn build_transport_hook_url(
|
||||
params.request_query,
|
||||
);
|
||||
}
|
||||
if let Some(auth_config) = resolve_local_vertex_service_account_auth_config(transport) {
|
||||
return build_vertex_service_account_gemini_content_url(
|
||||
params.mapped_model?,
|
||||
params.upstream_is_stream,
|
||||
&auth_config,
|
||||
params.request_query,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if is_antigravity_provider_transport(transport) {
|
||||
@@ -526,6 +535,43 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uses_vertex_service_account_hook_before_default_gemini_url() {
|
||||
let mut transport = sample_transport(
|
||||
"vertex_ai",
|
||||
"gemini:generate_content",
|
||||
"https://aiplatform.googleapis.com",
|
||||
None,
|
||||
);
|
||||
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":"svc@example.iam.gserviceaccount.com",
|
||||
"private_key":"TEST-PRIVATE-KEY",
|
||||
"project_id":"demo-project"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let url = build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "gemini:generate_content",
|
||||
mapped_model: Some("gemini-3.1-pro-preview"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("foo=bar&beta=1"),
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.expect("vertex service account hook url");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://aiplatform.googleapis.com/v1/projects/demo-project/locations/global/publishers/google/models/gemini-3.1-pro-preview:generateContent?foo=bar"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_openai_responses_url_for_formal_format_name() {
|
||||
let transport = sample_transport(
|
||||
|
||||
@@ -23,8 +23,8 @@ use crate::rules::{
|
||||
};
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
use crate::vertex::{
|
||||
is_vertex_api_key_transport_context,
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network,
|
||||
is_vertex_service_account_transport_context, is_vertex_transport_context,
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network,
|
||||
};
|
||||
use crate::{build_transport_request_url, ensure_upstream_auth_header, TransportRequestUrlParams};
|
||||
|
||||
@@ -103,7 +103,7 @@ pub fn classify_same_format_provider_request_behavior(
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("claude_code");
|
||||
let is_vertex = is_vertex_api_key_transport_context(transport);
|
||||
let is_vertex = is_vertex_transport_context(transport);
|
||||
let is_kiro = is_kiro_provider_transport(transport);
|
||||
let upstream_is_stream = aether_ai_formats::resolve_upstream_is_stream_from_endpoint_config(
|
||||
transport.endpoint.config.as_ref(),
|
||||
@@ -328,7 +328,7 @@ pub fn same_format_provider_transport_unsupported_reason(
|
||||
} else if behavior.is_claude_code {
|
||||
local_claude_code_transport_unsupported_reason_with_network(transport, api_format)
|
||||
} else if behavior.is_vertex {
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network(transport)
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network(transport)
|
||||
} else {
|
||||
match family {
|
||||
SameFormatProviderFamily::Standard => {
|
||||
@@ -391,6 +391,9 @@ pub fn should_try_same_format_provider_oauth_auth(
|
||||
behavior.is_kiro
|
||||
|| matches!(family, SameFormatProviderFamily::Standard)
|
||||
&& resolve_local_standard_auth(transport).is_none()
|
||||
|| matches!(family, SameFormatProviderFamily::Gemini)
|
||||
&& behavior.is_vertex
|
||||
&& is_vertex_service_account_transport_context(transport)
|
||||
|| matches!(family, SameFormatProviderFamily::Gemini)
|
||||
&& !behavior.is_vertex
|
||||
&& resolve_local_gemini_auth(transport).is_none()
|
||||
|
||||
@@ -1,6 +1,29 @@
|
||||
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 {
|
||||
@@ -8,6 +31,16 @@ pub struct VertexApiKeyQueryAuth {
|
||||
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> {
|
||||
@@ -39,14 +72,270 @@ pub fn resolve_local_vertex_api_key_query_auth(
|
||||
})
|
||||
}
|
||||
|
||||
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::{resolve_local_vertex_api_key_query_auth, VERTEX_API_KEY_QUERY_PARAM};
|
||||
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 {
|
||||
@@ -136,4 +425,60 @@ mod tests {
|
||||
.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":"svc@example.iam.gserviceaccount.com",
|
||||
"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, "svc@example.iam.gserviceaccount.com");
|
||||
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":"svc@example.iam.gserviceaccount.com",
|
||||
"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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -27,28 +27,36 @@ pub fn looks_like_vertex_ai_host(base_url: &str) -> bool {
|
||||
}
|
||||
|
||||
pub fn is_vertex_api_key_transport_context(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
if transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(super::PROVIDER_TYPE)
|
||||
{
|
||||
if is_vertex_provider_type(transport) {
|
||||
return resolve_local_auth_type_for_transport_format(transport)
|
||||
.eq_ignore_ascii_case("api_key");
|
||||
}
|
||||
|
||||
if !looks_like_vertex_ai_host(&transport.endpoint.base_url) {
|
||||
return false;
|
||||
}
|
||||
|
||||
let endpoint_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
|
||||
if !endpoint_api_format.starts_with("gemini:") && !endpoint_api_format.starts_with("claude:") {
|
||||
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,
|
||||
@@ -60,11 +68,28 @@ pub fn uses_vertex_api_key_query_auth(
|
||||
.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:")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
is_vertex_api_key_transport_context, looks_like_vertex_ai_host,
|
||||
uses_vertex_api_key_query_auth,
|
||||
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,
|
||||
@@ -150,6 +175,17 @@ mod tests {
|
||||
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();
|
||||
|
||||
@@ -4,20 +4,28 @@ mod policy;
|
||||
mod url;
|
||||
|
||||
pub use auth::{
|
||||
resolve_local_vertex_api_key_query_auth, VertexApiKeyQueryAuth, VERTEX_API_KEY_QUERY_PARAM,
|
||||
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, looks_like_vertex_ai_host, uses_vertex_api_key_query_auth,
|
||||
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_imagen_content_url,
|
||||
build_vertex_service_account_gemini_content_url, resolve_vertex_service_account_region,
|
||||
VERTEX_API_KEY_BASE_URL,
|
||||
};
|
||||
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::super::{
|
||||
body_rules_are_locally_supported, header_rules_are_locally_supported,
|
||||
resolve_transport_profile, transport_profile_is_configured,
|
||||
transport_proxy_is_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,
|
||||
};
|
||||
use super::auth::resolve_local_vertex_api_key_query_auth;
|
||||
|
||||
fn is_vertex_transport_family(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
@@ -17,6 +19,19 @@ fn is_vertex_transport_family(transport: &GatewayProviderTransportSnapshot) -> b
|
||||
|
||||
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 {
|
||||
@@ -41,7 +56,14 @@ pub fn local_vertex_api_key_gemini_transport_unsupported_reason_with_network(
|
||||
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
|
||||
return Some("transport_body_rules_unsupported");
|
||||
}
|
||||
if resolve_local_vertex_api_key_query_auth(transport).is_none() {
|
||||
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) {
|
||||
@@ -71,6 +93,12 @@ pub fn supports_local_vertex_api_key_gemini_transport_with_network(
|
||||
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 {
|
||||
@@ -159,8 +187,10 @@ mod tests {
|
||||
|
||||
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 {
|
||||
@@ -242,13 +272,33 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_vertex_service_account_subset() {
|
||||
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":"svc@example.iam.gserviceaccount.com",
|
||||
"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 allows_network_passthrough_for_custom_path_with_local_proxy_support() {
|
||||
let mut transport = sample_transport();
|
||||
@@ -272,5 +322,9 @@ mod tests {
|
||||
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")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ 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";
|
||||
|
||||
@@ -24,6 +25,15 @@ pub fn build_vertex_api_key_imagen_content_url(
|
||||
build_vertex_api_key_google_model_url(model, stream, 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> {
|
||||
build_vertex_service_account_google_model_url(model, stream, auth_config, request_query)
|
||||
}
|
||||
|
||||
fn build_vertex_api_key_google_model_url(
|
||||
model: &str,
|
||||
stream: bool,
|
||||
@@ -46,6 +56,89 @@ fn build_vertex_api_key_google_model_url(
|
||||
build_passthrough_path_url(VERTEX_API_KEY_BASE_URL, &path, merged_query.as_deref(), &[])
|
||||
}
|
||||
|
||||
fn build_vertex_service_account_google_model_url(
|
||||
model: &str,
|
||||
stream: bool,
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let trimmed_model = model.trim();
|
||||
let project_id = auth_config.project_id.trim();
|
||||
if trimmed_model.is_empty() || project_id.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let region = resolve_vertex_service_account_region(trimmed_model, auth_config);
|
||||
let action = if stream {
|
||||
"streamGenerateContent"
|
||||
} else {
|
||||
"generateContent"
|
||||
};
|
||||
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}:{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>,
|
||||
@@ -73,6 +166,29 @@ fn build_vertex_api_key_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;
|
||||
@@ -85,7 +201,13 @@ fn merge_query_string(out: &mut BTreeMap<String, String>, query: Option<&str>) {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{build_vertex_api_key_gemini_content_url, build_vertex_api_key_imagen_content_url};
|
||||
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() {
|
||||
@@ -118,4 +240,57 @@ mod tests {
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_vertex_service_account_gemini_sync_url() {
|
||||
let auth_config = VertexServiceAccountAuthConfig {
|
||||
client_email: "svc@example.iam.gserviceaccount.com".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: "svc@example.iam.gserviceaccount.com".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