fix(gateway): normalize Gemini Vertex embedding transport

This commit is contained in:
MMEXA
2026-05-18 15:39:29 +00:00
parent 66f154a251
commit b480f3aaff
24 changed files with 1046 additions and 107 deletions

View File

@@ -736,7 +736,7 @@ const ADMIN_API_FORMAT_DEFINITIONS: &[AdminApiFormatDefinition] = &[
AdminApiFormatDefinition {
value: "gemini:embedding",
label: "Gemini Embedding",
default_path: "/v1/embeddings",
default_path: "/v1beta/models/{model}:{action}",
aliases: &["gemini_embedding"],
},
AdminApiFormatDefinition {

View File

@@ -11,7 +11,9 @@ pub fn from(body_json: &Value) -> Option<CanonicalEmbeddingResponse> {
return None;
}
let embeddings = if let Some(values) = body
let embeddings = if let Some(raw_embeddings) = vertex_predict_embeddings(body) {
raw_embeddings
} else if let Some(values) = body
.get("embedding")
.and_then(Value::as_object)
.and_then(|embedding| embedding.get("values"))
@@ -56,6 +58,7 @@ pub fn from(body_json: &Value) -> Option<CanonicalEmbeddingResponse> {
model: body
.get("model")
.or_else(|| body.get("modelVersion"))
.or_else(|| body.get("deployedModelId"))
.and_then(Value::as_str)
.unwrap_or("unknown")
.to_string(),
@@ -69,14 +72,36 @@ pub fn from(body_json: &Value) -> Option<CanonicalEmbeddingResponse> {
"responseId",
"model",
"modelVersion",
"deployedModelId",
"embedding",
"embeddings",
"predictions",
"usageMetadata",
],
),
})
}
fn vertex_predict_embeddings(
body: &serde_json::Map<String, Value>,
) -> Option<Vec<CanonicalEmbedding>> {
let predictions = body.get("predictions")?.as_array()?;
predictions
.iter()
.enumerate()
.map(|(index, item)| {
let item_object = item.as_object()?;
let embedding_object = item_object.get("embeddings")?.as_object()?;
let values = embedding_object.get("values")?.as_array()?;
Some(CanonicalEmbedding {
index,
embedding: embedding_values(values)?,
extensions: namespace_extensions("vertex", item_object, &["embeddings"]),
})
})
.collect()
}
fn embedding_values(values: &[Value]) -> Option<Vec<f64>> {
values.iter().map(Value::as_f64).collect()
}
@@ -105,4 +130,29 @@ mod tests {
assert_eq!(usage.input_tokens, 4);
assert_eq!(usage.total_tokens, 4);
}
#[test]
fn parses_vertex_predict_embedding_response() {
let body = json!({
"predictions": [
{
"embeddings": {
"values": [0.1, 0.2, 0.3]
}
},
{
"embeddings": {
"values": [0.4, 0.5, 0.6]
}
}
],
"deployedModelId": "gemini-embedding-2"
});
let parsed = from(&body).expect("response should parse");
assert_eq!(parsed.model, "gemini-embedding-2");
assert_eq!(parsed.embeddings[0].embedding, vec![0.1, 0.2, 0.3]);
assert_eq!(parsed.embeddings[1].embedding, vec![0.4, 0.5, 0.6]);
}
}

View File

@@ -162,6 +162,21 @@ impl CandidateFailureDiagnostic {
.source(source)
}
pub fn request_conversion_failed(
client_api_format: impl Into<String>,
provider_api_format: impl Into<String>,
source: impl Into<String>,
message: impl Into<String>,
) -> Self {
Self::new(
CandidateFailureDiagnosticKind::RequestConversion,
"$",
message,
)
.formats(client_api_format, provider_api_format)
.source(source)
}
pub fn envelope_build_failed(
client_api_format: impl Into<String>,
provider_api_format: impl Into<String>,

View File

@@ -14,6 +14,7 @@ pub mod oauth_refresh;
mod openai_image;
pub mod policy;
pub mod provider_types;
mod request_body;
mod request_url;
pub mod rules;
pub mod same_format_provider;
@@ -69,6 +70,9 @@ pub use policy::{
local_standard_transport_unsupported_reason_with_network, supports_local_gemini_transport,
supports_local_gemini_transport_with_network, supports_local_standard_transport,
};
pub use request_body::{
apply_transport_request_body_semantics, TransportRequestBodySemanticsError,
};
pub use request_url::{
build_cross_format_openai_chat_upstream_url, build_cross_format_openai_responses_upstream_url,
build_kiro_cross_format_upstream_url, build_local_openai_chat_upstream_url,

View File

@@ -0,0 +1,415 @@
use serde_json::{Map, Value};
use crate::snapshot::GatewayProviderTransportSnapshot;
use crate::vertex::is_vertex_transport_context;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TransportRequestBodySemanticsError {
message: &'static str,
}
impl TransportRequestBodySemanticsError {
const fn new(message: &'static str) -> Self {
Self { message }
}
pub const fn message(&self) -> &'static str {
self.message
}
}
impl std::fmt::Display for TransportRequestBodySemanticsError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.message)
}
}
impl std::error::Error for TransportRequestBodySemanticsError {}
pub fn apply_transport_request_body_semantics(
provider_request_body: &mut Value,
transport: &GatewayProviderTransportSnapshot,
provider_api_format: &str,
) -> Result<(), TransportRequestBodySemanticsError> {
let provider_api_format = aether_ai_formats::normalize_api_format_alias(provider_api_format);
if provider_api_format == "gemini:embedding" && is_vertex_transport_context(transport) {
apply_vertex_gemini_embedding_body_semantics(provider_request_body)?;
}
Ok(())
}
fn apply_vertex_gemini_embedding_body_semantics(
provider_request_body: &mut Value,
) -> Result<(), TransportRequestBodySemanticsError> {
let object = provider_request_body.as_object_mut().ok_or_else(|| {
TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding request body must be a JSON object",
)
})?;
if object.contains_key("instances") {
validate_existing_vertex_predict_body(object)?;
object.remove("model");
return Ok(());
}
let next = build_vertex_predict_body_from_gemini_embedding_object(object)?;
*object = next;
Ok(())
}
fn build_vertex_predict_body_from_gemini_embedding_object(
object: &Map<String, Value>,
) -> Result<Map<String, Value>, TransportRequestBodySemanticsError> {
if let Some(requests) = object.get("requests") {
if object.keys().any(|key| key != "requests") {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding batch body cannot mix requests with other top-level fields",
));
}
let request_items = requests.as_array().ok_or_else(|| {
TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding requests must be an array",
)
})?;
let request_objects = request_items
.iter()
.map(Value::as_object)
.collect::<Option<Vec<_>>>()
.ok_or_else(|| {
TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding requests must be an array of objects",
)
})?;
if request_objects.is_empty() {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding requests must contain at least one item",
));
}
return build_vertex_predict_body_from_gemini_embedding_items(&request_objects);
}
build_vertex_predict_body_from_gemini_embedding_items(&[object])
}
fn build_vertex_predict_body_from_gemini_embedding_items(
items: &[&Map<String, Value>],
) -> Result<Map<String, Value>, TransportRequestBodySemanticsError> {
if items.iter().any(|item| {
item.keys().any(|key| {
!matches!(
key.as_str(),
"model"
| "content"
| "taskType"
| "title"
| "outputDimensionality"
| "autoTruncate"
)
})
}) {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding body contains fields that cannot be mapped to predict instances",
));
}
let instances = items
.iter()
.map(|item| build_vertex_predict_instance(item))
.collect::<Option<Vec<_>>>()
.ok_or_else(|| {
TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding body must contain text content parts",
)
})?;
if instances.is_empty() {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding body must contain at least one instance",
));
}
let mut output = Map::new();
output.insert("instances".to_string(), Value::Array(instances));
let mut parameters = Map::new();
insert_shared_parameter(items, &mut parameters, "outputDimensionality")?;
insert_shared_parameter(items, &mut parameters, "autoTruncate")?;
if !parameters.is_empty() {
output.insert("parameters".to_string(), Value::Object(parameters));
}
Ok(output)
}
fn validate_existing_vertex_predict_body(
object: &Map<String, Value>,
) -> Result<(), TransportRequestBodySemanticsError> {
if object
.keys()
.any(|key| !matches!(key.as_str(), "model" | "instances" | "parameters"))
{
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding predict body contains unsupported top-level fields",
));
}
let Some(instances) = object.get("instances").and_then(Value::as_array) else {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding predict body must contain an instances array",
));
};
if instances.is_empty() {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding predict body must contain at least one instance",
));
}
if object
.get("parameters")
.is_some_and(|parameters| !parameters.is_object())
{
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding predict parameters must be an object",
));
}
Ok(())
}
fn build_vertex_predict_instance(item: &Map<String, Value>) -> Option<Value> {
let content = gemini_embedding_content_text(item.get("content")?)?;
let mut instance = Map::new();
instance.insert("content".to_string(), Value::String(content));
if let Some(task_type) = item.get("taskType") {
instance.insert(
"task_type".to_string(),
Value::String(task_type.as_str()?.to_string()),
);
}
if let Some(title) = item.get("title") {
instance.insert(
"title".to_string(),
Value::String(title.as_str()?.to_string()),
);
}
Some(Value::Object(instance))
}
fn gemini_embedding_content_text(content: &Value) -> Option<String> {
let parts = content
.as_object()?
.get("parts")?
.as_array()?
.iter()
.filter_map(|part| part.as_object()?.get("text")?.as_str())
.filter(|text| !text.trim().is_empty())
.collect::<Vec<_>>();
if parts.is_empty() {
return None;
}
Some(parts.join(""))
}
fn insert_shared_parameter(
items: &[&Map<String, Value>],
parameters: &mut Map<String, Value>,
key: &str,
) -> Result<(), TransportRequestBodySemanticsError> {
let mut value: Option<Value> = None;
for item in items {
let Some(next) = item.get(key) else {
continue;
};
match &value {
Some(current) if current != next => {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding batch items must use the same shared parameters",
));
}
None => value = Some(next.clone()),
_ => {}
}
}
if let Some(value) = value {
parameters.insert(key.to_string(), value);
}
Ok(())
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::apply_transport_request_body_semantics;
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport(provider_type: &str, base_url: &str) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "provider".to_string(),
provider_type: provider_type.to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
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:embedding".to_string(),
api_family: Some("gemini".to_string()),
endpoint_kind: Some("embedding".to_string()),
is_active: true,
base_url: base_url.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:embedding".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: "secret".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn vertex_gemini_embedding_single_body_uses_predict_contract() {
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
let mut body = json!({
"model": "gemini-embedding-2",
"content": {"parts": [{"text": "hello"}]},
"taskType": "RETRIEVAL_QUERY",
"outputDimensionality": 768
});
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect("body semantics should apply");
assert!(body.get("model").is_none());
assert!(body.get("content").is_none());
assert_eq!(body["instances"][0]["content"], "hello");
assert_eq!(body["instances"][0]["task_type"], "RETRIEVAL_QUERY");
assert_eq!(body["parameters"]["outputDimensionality"], 768);
}
#[test]
fn gemini_api_embedding_single_body_keeps_model_for_developer_api() {
let transport =
sample_transport("gemini", "https://generativelanguage.googleapis.com/v1beta");
let mut body = json!({
"model": "gemini-embedding-2",
"content": {"parts": [{"text": "hello"}]}
});
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect("developer API body should pass through");
assert_eq!(body["model"], "gemini-embedding-2");
}
#[test]
fn vertex_gemini_embedding_batch_body_uses_predict_instances() {
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
let mut body = json!({
"requests": [
{
"model": "models/gemini-embedding-2",
"content": {"parts": [{"text": "hello"}]}
},
{
"model": "models/gemini-embedding-2",
"content": {"parts": [{"text": "world"}]}
}
]
});
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect("batch body semantics should apply");
assert!(body.get("requests").is_none());
assert_eq!(body["instances"][0]["content"], "hello");
assert_eq!(body["instances"][1]["content"], "world");
}
#[test]
fn vertex_gemini_embedding_existing_predict_body_removes_duplicate_model() {
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
let mut body = json!({
"model": "gemini-embedding-2",
"instances": [
{"content": "hello", "task_type": "RETRIEVAL_QUERY"}
],
"parameters": {
"outputDimensionality": 768
}
});
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect("existing predict body should be accepted");
assert!(body.get("model").is_none());
assert_eq!(body["instances"][0]["content"], "hello");
assert_eq!(body["parameters"]["outputDimensionality"], 768);
}
#[test]
fn vertex_gemini_embedding_existing_predict_body_rejects_unconsumed_fields() {
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
let mut body = json!({
"model": "gemini-embedding-2",
"instances": [
{"content": "hello"}
],
"input": "this field would not be consumed by Vertex predict"
});
let error =
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect_err("predict body must not carry unconsumed OpenAI fields");
assert!(error.message().contains("unsupported top-level fields"));
assert!(body.get("model").is_some());
}
#[test]
fn vertex_gemini_embedding_rejects_unconverted_openai_body() {
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
let mut body = json!({
"model": "gemini-embedding-2",
"input": "hello"
});
let error =
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect_err("OpenAI embedding body must not be sent to Vertex native predict");
assert!(error.message().contains("cannot be mapped"));
assert!(body.get("input").is_some());
}
}

View File

@@ -19,8 +19,8 @@ use crate::url::{
use crate::vertex::{
build_vertex_api_key_gemini_content_url, build_vertex_api_key_gemini_embedding_url,
build_vertex_service_account_gemini_content_url,
build_vertex_service_account_gemini_embedding_url, is_vertex_transport_context,
resolve_local_vertex_api_key_query_auth, resolve_local_vertex_service_account_auth_config,
build_vertex_service_account_gemini_embedding_url, resolve_local_vertex_api_key_query_auth,
resolve_local_vertex_service_account_auth_config,
};
#[derive(Debug, Clone, Copy)]
@@ -68,13 +68,6 @@ fn build_transport_request_url_inner(
let provider_api_format = params.provider_api_format.trim().to_ascii_lowercase();
let normalized_provider_api_format =
aether_ai_formats::normalize_api_format_alias(&provider_api_format);
if normalized_provider_api_format == "gemini:embedding"
&& gemini_embedding_batch
&& is_vertex_transport_context(transport)
{
return None;
}
if let Some(url) = build_transport_hook_url(transport, params) {
return Some(url);
}
@@ -696,12 +689,12 @@ mod tests {
assert_eq!(
url,
"https://aiplatform.googleapis.com/v1/projects/demo-project/locations/global/publishers/google/models/gemini-embedding-2:embedContent?foo=bar"
"https://aiplatform.googleapis.com/v1/projects/demo-project/locations/global/publishers/google/models/gemini-embedding-2:predict?foo=bar"
);
}
#[test]
fn vertex_gemini_embedding_batch_request_does_not_use_gemini_api_batch_endpoint() {
fn vertex_gemini_embedding_batch_request_uses_vertex_predict_endpoint() {
let mut transport = sample_transport(
"vertex_ai",
"gemini:embedding",
@@ -729,18 +722,23 @@ mod tests {
]
});
assert!(build_transport_request_url_for_request_body(
&transport,
TransportRequestUrlParams {
provider_api_format: "gemini:embedding",
mapped_model: Some("gemini-embedding-2"),
upstream_is_stream: false,
request_query: None,
kiro_api_region: None,
},
Some(&batch_body),
)
.is_none());
assert_eq!(
build_transport_request_url_for_request_body(
&transport,
TransportRequestUrlParams {
provider_api_format: "gemini:embedding",
mapped_model: Some("gemini-embedding-2"),
upstream_is_stream: false,
request_query: None,
kiro_api_region: None,
},
Some(&batch_body),
)
.as_deref(),
Some(
"https://aiplatform.googleapis.com/v1/projects/demo-project/locations/global/publishers/google/models/gemini-embedding-2:predict"
)
);
}
#[test]

View File

@@ -40,7 +40,7 @@ pub fn build_vertex_api_key_gemini_embedding_url(
api_key: &str,
request_query: Option<&str>,
) -> Option<String> {
build_vertex_api_key_google_model_url(model, "embedContent", false, api_key, request_query)
build_vertex_api_key_google_model_url(model, "predict", false, api_key, request_query)
}
pub fn build_vertex_service_account_gemini_content_url(
@@ -64,7 +64,7 @@ pub fn build_vertex_service_account_gemini_embedding_url(
) -> Option<String> {
build_vertex_service_account_google_model_url(
model,
"embedContent",
"predict",
false,
auth_config,
request_query,