mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
fix(gateway): normalize Gemini Vertex embedding transport
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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>,
|
||||
|
||||
@@ -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,
|
||||
|
||||
415
crates/aether-provider-transport/src/request_body.rs
Normal file
415
crates/aether-provider-transport/src/request_body.rs
Normal 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());
|
||||
}
|
||||
}
|
||||
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user