mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
feat: add embedding and rerank support
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
use super::provider_types::{
|
||||
provider_type_supports_local_embedding_transport,
|
||||
provider_type_supports_local_openai_chat_transport,
|
||||
provider_type_supports_local_same_format_transport,
|
||||
};
|
||||
@@ -195,10 +196,200 @@ fn local_same_format_transport_unsupported_reason(
|
||||
if !provider_type_supported(&transport.provider.provider_type) {
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
if aether_ai_formats::is_embedding_api_format(api_format) {
|
||||
if !endpoint_kind_allows_embedding(transport.endpoint.endpoint_kind.as_deref()) {
|
||||
return Some("transport_endpoint_kind_unsupported");
|
||||
}
|
||||
if !provider_type_supports_local_embedding_transport(
|
||||
&transport.provider.provider_type,
|
||||
api_format,
|
||||
) {
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
}
|
||||
if aether_ai_formats::is_rerank_api_format(api_format) {
|
||||
if !endpoint_kind_allows_rerank(transport.endpoint.endpoint_kind.as_deref()) {
|
||||
return Some("transport_endpoint_kind_unsupported");
|
||||
}
|
||||
if !provider_type_supports_local_embedding_transport(
|
||||
&transport.provider.provider_type,
|
||||
api_format,
|
||||
) {
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
fn endpoint_kind_allows_embedding(endpoint_kind: Option<&str>) -> bool {
|
||||
endpoint_kind
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| {
|
||||
matches!(
|
||||
value.to_ascii_lowercase().as_str(),
|
||||
"embedding" | "embeddings"
|
||||
)
|
||||
})
|
||||
.unwrap_or(true)
|
||||
}
|
||||
|
||||
fn endpoint_kind_allows_rerank(endpoint_kind: Option<&str>) -> bool {
|
||||
endpoint_kind
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| matches!(value.to_ascii_lowercase().as_str(), "rerank" | "reranking"))
|
||||
.unwrap_or(true)
|
||||
}
|
||||
|
||||
fn same_api_format(left: &str, right: &str) -> bool {
|
||||
aether_ai_formats::api_format_alias_matches(left, right)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::local_standard_transport_unsupported_reason_with_network;
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport(
|
||||
provider_type: &str,
|
||||
api_format: &str,
|
||||
endpoint_kind: Option<&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: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: api_format.to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: endpoint_kind.map(ToOwned::to_owned),
|
||||
is_active: true,
|
||||
base_url: "https://provider.example".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: None,
|
||||
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: "sk-test".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_unsupported_embedding_provider_format_pairs() {
|
||||
let openai_on_gemini = sample_transport("openai", "gemini:embedding", Some("embedding"));
|
||||
let gemini_on_openai = sample_transport("gemini", "openai:embedding", Some("embedding"));
|
||||
let chat_marked_embedding = sample_transport("openai", "openai:embedding", Some("chat"));
|
||||
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
&openai_on_gemini,
|
||||
"gemini:embedding"
|
||||
),
|
||||
Some("transport_provider_type_unsupported")
|
||||
);
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
&gemini_on_openai,
|
||||
"openai:embedding"
|
||||
),
|
||||
Some("transport_provider_type_unsupported")
|
||||
);
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
&chat_marked_embedding,
|
||||
"openai:embedding"
|
||||
),
|
||||
Some("transport_endpoint_kind_unsupported")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_supported_embedding_provider_format_pairs() {
|
||||
for (provider_type, api_format) in [
|
||||
("openai", "openai:embedding"),
|
||||
("gemini", "gemini:embedding"),
|
||||
("google", "gemini:embedding"),
|
||||
("jina", "jina:embedding"),
|
||||
("doubao", "doubao:embedding"),
|
||||
("volcengine", "doubao:embedding"),
|
||||
("custom", "openai:embedding"),
|
||||
("custom", "gemini:embedding"),
|
||||
("custom", "jina:embedding"),
|
||||
("custom", "doubao:embedding"),
|
||||
] {
|
||||
let transport = sample_transport(provider_type, api_format, Some("embedding"));
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(&transport, api_format),
|
||||
None,
|
||||
"{provider_type} should support {api_format}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_policy_accepts_embedding_endpoint_kind_aliases_only() {
|
||||
for endpoint_kind in [None, Some(""), Some(" embedding "), Some("EMBEDDINGS")] {
|
||||
let transport = sample_transport("openai", "openai:embedding", endpoint_kind);
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
&transport,
|
||||
"openai:embedding"
|
||||
),
|
||||
None,
|
||||
"endpoint kind {endpoint_kind:?} should be accepted"
|
||||
);
|
||||
}
|
||||
|
||||
for endpoint_kind in [Some("chat"), Some("responses"), Some("image")] {
|
||||
let transport = sample_transport("openai", "openai:embedding", endpoint_kind);
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
&transport,
|
||||
"openai:embedding"
|
||||
),
|
||||
Some("transport_endpoint_kind_unsupported"),
|
||||
"endpoint kind {endpoint_kind:?} should fail closed"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -225,6 +225,24 @@ pub fn provider_type_supports_local_same_format_transport(provider_type: &str) -
|
||||
)
|
||||
}
|
||||
|
||||
pub fn provider_type_supports_local_embedding_transport(
|
||||
provider_type: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
let provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
let api_format = aether_ai_formats::normalize_api_format_alias(api_format);
|
||||
|
||||
match api_format.as_str() {
|
||||
"openai:embedding" => matches!(provider_type.as_str(), "custom" | "openai"),
|
||||
"openai:rerank" => matches!(provider_type.as_str(), "custom" | "openai"),
|
||||
"gemini:embedding" => matches!(provider_type.as_str(), "custom" | "gemini" | "google"),
|
||||
"jina:embedding" => matches!(provider_type.as_str(), "custom" | "jina"),
|
||||
"jina:rerank" => matches!(provider_type.as_str(), "custom" | "jina"),
|
||||
"doubao:embedding" => matches!(provider_type.as_str(), "custom" | "doubao" | "volcengine"),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_codex_cli_backend_url(url: &str) -> bool {
|
||||
let url = url.trim().to_ascii_lowercase();
|
||||
url.contains("/codex") && (url.contains("/backend-api/") || url.contains("/backendapi/"))
|
||||
@@ -301,7 +319,8 @@ pub const ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES: &[&str] =
|
||||
mod tests {
|
||||
use super::{
|
||||
fixed_provider_endpoint_template_by_api_format, fixed_provider_key_inherits_api_formats,
|
||||
fixed_provider_template, FixedProviderEndpointConfigValue,
|
||||
fixed_provider_template, provider_type_supports_local_embedding_transport,
|
||||
FixedProviderEndpointConfigValue,
|
||||
};
|
||||
|
||||
#[test]
|
||||
@@ -355,4 +374,41 @@ mod tests {
|
||||
"custom", "oauth", None
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_type_supports_only_matching_embedding_formats() {
|
||||
for (provider_type, api_format) in [
|
||||
("openai", "openai:embedding"),
|
||||
("custom", "openai:embedding"),
|
||||
("gemini", "gemini:embedding"),
|
||||
("google", "gemini:embedding"),
|
||||
("jina", "jina:embedding"),
|
||||
("doubao", "doubao:embedding"),
|
||||
("volcengine", "doubao:embedding"),
|
||||
] {
|
||||
assert!(
|
||||
provider_type_supports_local_embedding_transport(provider_type, api_format),
|
||||
"{provider_type} should support {api_format}"
|
||||
);
|
||||
}
|
||||
|
||||
for (provider_type, api_format) in [
|
||||
("openai", "gemini:embedding"),
|
||||
("gemini", "openai:embedding"),
|
||||
("jina", "doubao:embedding"),
|
||||
("doubao", "jina:embedding"),
|
||||
("claude_code", "openai:embedding"),
|
||||
("openai", "openai:chat"),
|
||||
] {
|
||||
assert!(
|
||||
!provider_type_supports_local_embedding_transport(provider_type, api_format),
|
||||
"{provider_type} should not support {api_format}"
|
||||
);
|
||||
}
|
||||
|
||||
assert!(provider_type_supports_local_embedding_transport(
|
||||
" Google ",
|
||||
"GEMINI:EMBEDDING"
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -78,6 +78,12 @@ pub fn build_transport_request_url(
|
||||
params.request_query,
|
||||
true,
|
||||
)),
|
||||
"openai:embedding" | "jina:embedding" => {
|
||||
build_provider_embedding_v1_url(&transport.endpoint.base_url, params.request_query)
|
||||
}
|
||||
"openai:rerank" | "jina:rerank" => {
|
||||
build_provider_rerank_v1_url(&transport.endpoint.base_url, params.request_query)
|
||||
}
|
||||
"claude:messages" => Some(build_claude_messages_url(
|
||||
&transport.endpoint.base_url,
|
||||
params.request_query,
|
||||
@@ -88,6 +94,17 @@ pub fn build_transport_request_url(
|
||||
params.upstream_is_stream,
|
||||
params.request_query,
|
||||
),
|
||||
"gemini:embedding" => build_gemini_embedding_url(
|
||||
&transport.endpoint.base_url,
|
||||
params.mapped_model?,
|
||||
params.request_query,
|
||||
),
|
||||
"doubao:embedding" => build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
"/embeddings/multimodal",
|
||||
params.request_query,
|
||||
&[],
|
||||
),
|
||||
_ => None,
|
||||
}?;
|
||||
|
||||
@@ -221,11 +238,8 @@ fn build_transport_hook_url(
|
||||
));
|
||||
}
|
||||
|
||||
if params
|
||||
.provider_api_format
|
||||
.trim()
|
||||
.to_ascii_lowercase()
|
||||
.starts_with("gemini:")
|
||||
if aether_ai_formats::normalize_api_format_alias(params.provider_api_format)
|
||||
== "gemini:generate_content"
|
||||
{
|
||||
if let Some(auth) = resolve_local_vertex_api_key_query_auth(transport) {
|
||||
return build_vertex_api_key_gemini_content_url(
|
||||
@@ -266,15 +280,14 @@ fn build_path_params(params: TransportRequestUrlParams<'_>) -> BTreeMap<&'static
|
||||
{
|
||||
path_params.insert("model", model);
|
||||
}
|
||||
if params
|
||||
.provider_api_format
|
||||
.trim()
|
||||
.to_ascii_lowercase()
|
||||
.starts_with("gemini:")
|
||||
{
|
||||
let provider_api_format =
|
||||
aether_ai_formats::normalize_api_format_alias(params.provider_api_format);
|
||||
if provider_api_format.starts_with("gemini:") {
|
||||
path_params.insert(
|
||||
"action",
|
||||
if params.upstream_is_stream {
|
||||
if provider_api_format == "gemini:embedding" {
|
||||
"embedContent"
|
||||
} else if params.upstream_is_stream {
|
||||
"streamGenerateContent"
|
||||
} else {
|
||||
"generateContent"
|
||||
@@ -284,6 +297,60 @@ fn build_path_params(params: TransportRequestUrlParams<'_>) -> BTreeMap<&'static
|
||||
path_params
|
||||
}
|
||||
|
||||
fn build_provider_embedding_v1_url(upstream_base_url: &str, query: Option<&str>) -> Option<String> {
|
||||
build_provider_v1_url(upstream_base_url, "/embeddings", "/v1/embeddings", query)
|
||||
}
|
||||
|
||||
fn build_provider_rerank_v1_url(upstream_base_url: &str, query: Option<&str>) -> Option<String> {
|
||||
build_provider_v1_url(upstream_base_url, "/rerank", "/v1/rerank", query)
|
||||
}
|
||||
|
||||
fn build_provider_v1_url(
|
||||
upstream_base_url: &str,
|
||||
v1_path: &str,
|
||||
default_path: &str,
|
||||
query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let base_without_query = upstream_base_url
|
||||
.trim()
|
||||
.split_once('?')
|
||||
.map(|(base, _)| base)
|
||||
.unwrap_or_else(|| upstream_base_url.trim())
|
||||
.trim_end_matches('/');
|
||||
let path = if base_without_query.ends_with("/v1") {
|
||||
v1_path
|
||||
} else {
|
||||
default_path
|
||||
};
|
||||
build_passthrough_path_url(upstream_base_url, path, query, &[])
|
||||
}
|
||||
|
||||
fn build_gemini_embedding_url(
|
||||
upstream_base_url: &str,
|
||||
model: &str,
|
||||
query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let trimmed_base_url = upstream_base_url
|
||||
.trim()
|
||||
.split_once('?')
|
||||
.map(|(base, _)| base)
|
||||
.unwrap_or_else(|| upstream_base_url.trim())
|
||||
.trim_end_matches('/');
|
||||
let trimmed_model = model.trim();
|
||||
if trimmed_base_url.is_empty() || trimmed_model.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let path = if trimmed_base_url.ends_with("/v1beta") {
|
||||
format!("/models/{trimmed_model}:embedContent")
|
||||
} else if trimmed_base_url.contains("/v1beta/models/") {
|
||||
":embedContent".to_string()
|
||||
} else {
|
||||
format!("/v1beta/models/{trimmed_model}:embedContent")
|
||||
};
|
||||
build_passthrough_path_url(upstream_base_url, &path, query, &["key"])
|
||||
}
|
||||
|
||||
fn expand_custom_path_template(path: &str, params: BTreeMap<&'static str, &str>) -> String {
|
||||
if params.is_empty() {
|
||||
return path.to_string();
|
||||
@@ -320,7 +387,10 @@ fn maybe_add_gemini_stream_alt_sse(
|
||||
provider_api_format: &str,
|
||||
upstream_is_stream: bool,
|
||||
) -> String {
|
||||
if !provider_api_format.starts_with("gemini:") || !upstream_is_stream {
|
||||
if aether_ai_formats::normalize_api_format_alias(provider_api_format)
|
||||
!= "gemini:generate_content"
|
||||
|| !upstream_is_stream
|
||||
{
|
||||
return upstream_url;
|
||||
}
|
||||
|
||||
@@ -548,4 +618,260 @@ mod tests {
|
||||
));
|
||||
assert!(url.contains("conversationId=abc"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_request_url_builds_provider_default_paths() {
|
||||
let openai = sample_transport(
|
||||
"openai",
|
||||
"openai:embedding",
|
||||
"https://api.openai.example/v1",
|
||||
None,
|
||||
);
|
||||
let jina = sample_transport("jina", "jina:embedding", "https://api.jina.example", None);
|
||||
let gemini = sample_transport(
|
||||
"gemini",
|
||||
"gemini:embedding",
|
||||
"https://generativelanguage.googleapis.com/v1beta",
|
||||
None,
|
||||
);
|
||||
let doubao = sample_transport(
|
||||
"doubao",
|
||||
"doubao:embedding",
|
||||
"https://ark.volces.example/api/v3",
|
||||
None,
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&openai,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "openai:embedding",
|
||||
mapped_model: Some("text-embedding-3-small"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("tenant=demo"),
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://api.openai.example/v1/embeddings?tenant=demo")
|
||||
);
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&jina,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "jina:embedding",
|
||||
mapped_model: Some("jina-embeddings-v3"),
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://api.jina.example/v1/embeddings")
|
||||
);
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&gemini,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "gemini:embedding",
|
||||
mapped_model: Some("gemini-embedding-001"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-key&foo=bar"),
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-001:embedContent?foo=bar"
|
||||
)
|
||||
);
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&doubao,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "doubao:embedding",
|
||||
mapped_model: Some("doubao-embedding-vision"),
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://ark.volces.example/api/v3/embeddings/multimodal")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rerank_request_url_builds_provider_default_paths() {
|
||||
let openai = sample_transport(
|
||||
"openai",
|
||||
"openai:rerank",
|
||||
"https://api.openai.example/v1",
|
||||
None,
|
||||
);
|
||||
let jina = sample_transport("jina", "jina:rerank", "https://api.jina.example", None);
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&openai,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "openai:rerank",
|
||||
mapped_model: Some("rerank-1"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("tenant=demo"),
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://api.openai.example/v1/rerank?tenant=demo")
|
||||
);
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&jina,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "jina:rerank",
|
||||
mapped_model: Some("jina-reranker-v2-base-multilingual"),
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://api.jina.example/v1/rerank")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_request_url_handles_base_variants_and_queries() {
|
||||
let openai_without_v1 = sample_transport(
|
||||
"openai",
|
||||
"openai:embedding",
|
||||
"https://api.openai.example/root?tenant=base",
|
||||
None,
|
||||
);
|
||||
let jina_with_v1 = sample_transport(
|
||||
"jina",
|
||||
"jina:embedding",
|
||||
"https://api.jina.example/v1/",
|
||||
None,
|
||||
);
|
||||
let gemini_model_base = sample_transport(
|
||||
"gemini",
|
||||
"gemini:embedding",
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-001",
|
||||
None,
|
||||
);
|
||||
let doubao_with_query = sample_transport(
|
||||
"doubao",
|
||||
"doubao:embedding",
|
||||
"https://ark.volces.example/api/v3?tenant=base",
|
||||
None,
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&openai_without_v1,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "OPENAI:EMBEDDING",
|
||||
mapped_model: Some("text-embedding-3-small"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("tenant=request&trace=1"),
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://api.openai.example/root/v1/embeddings?tenant=request&trace=1")
|
||||
);
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&jina_with_v1,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "jina:embedding",
|
||||
mapped_model: Some("jina-embeddings-v3"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("trace=2"),
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://api.jina.example/v1/embeddings?trace=2")
|
||||
);
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&gemini_model_base,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "gemini:embedding",
|
||||
mapped_model: Some("gemini-embedding-001"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-key&trace=3"),
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-001:embedContent?trace=3")
|
||||
);
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&doubao_with_query,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "doubao:embedding",
|
||||
mapped_model: Some("doubao-embedding-vision"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("trace=4"),
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://ark.volces.example/api/v3/embeddings/multimodal?tenant=base&trace=4")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gemini_embedding_request_url_requires_mapped_model_without_custom_path() {
|
||||
let transport = sample_transport(
|
||||
"gemini",
|
||||
"gemini:embedding",
|
||||
"https://generativelanguage.googleapis.com/v1beta",
|
||||
None,
|
||||
);
|
||||
|
||||
assert!(build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "gemini:embedding",
|
||||
mapped_model: None,
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_request_url_expands_custom_gemini_embed_action() {
|
||||
let transport = sample_transport(
|
||||
"custom",
|
||||
"gemini:embedding",
|
||||
"https://generativelanguage.googleapis.com",
|
||||
Some("/v1beta/models/{model}:{action}"),
|
||||
);
|
||||
|
||||
let url = build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "gemini:embedding",
|
||||
mapped_model: Some("gemini-embedding-001"),
|
||||
upstream_is_stream: true,
|
||||
request_query: Some("key=client-key&foo=bar"),
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
.expect("expanded custom embedding path url");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-001:embedContent?foo=bar"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -52,6 +52,7 @@ pub struct SameFormatProviderRequestBehavior {
|
||||
pub struct SameFormatProviderRequestBodyInput<'a> {
|
||||
pub body_json: &'a Value,
|
||||
pub mapped_model: &'a str,
|
||||
pub client_api_format: &'a str,
|
||||
pub provider_api_format: &'a str,
|
||||
pub source_model: Option<&'a str>,
|
||||
pub family: SameFormatProviderFamily,
|
||||
@@ -133,12 +134,27 @@ pub fn build_same_format_provider_request_body(
|
||||
);
|
||||
}
|
||||
|
||||
let request_body_object = input.body_json.as_object()?;
|
||||
let mut provider_request_body = serde_json::Map::from_iter(
|
||||
request_body_object
|
||||
.iter()
|
||||
.map(|(key, value)| (key.clone(), value.clone())),
|
||||
);
|
||||
let mut provider_request_body = if aether_ai_formats::api_format_alias_matches(
|
||||
input.client_api_format,
|
||||
input.provider_api_format,
|
||||
) {
|
||||
let request_body_object = input.body_json.as_object()?;
|
||||
serde_json::Map::from_iter(
|
||||
request_body_object
|
||||
.iter()
|
||||
.map(|(key, value)| (key.clone(), value.clone())),
|
||||
)
|
||||
} else {
|
||||
aether_ai_formats::convert_request(
|
||||
input.client_api_format,
|
||||
input.provider_api_format,
|
||||
input.body_json,
|
||||
&aether_ai_formats::FormatContext::default().with_mapped_model(input.mapped_model),
|
||||
)
|
||||
.ok()?
|
||||
.as_object()?
|
||||
.clone()
|
||||
};
|
||||
match input.family {
|
||||
SameFormatProviderFamily::Standard => {
|
||||
provider_request_body.insert(
|
||||
@@ -488,6 +504,7 @@ mod tests {
|
||||
"messages": [{"role": "user", "content": "hello"}]
|
||||
}),
|
||||
mapped_model: "upstream-model",
|
||||
client_api_format: "openai:chat",
|
||||
provider_api_format: "openai:chat",
|
||||
source_model: Some("client-model"),
|
||||
family: SameFormatProviderFamily::Standard,
|
||||
@@ -512,6 +529,7 @@ mod tests {
|
||||
"reasoning_effort": "low"
|
||||
}),
|
||||
mapped_model: "upstream-model",
|
||||
client_api_format: "openai:chat",
|
||||
provider_api_format: "openai:chat",
|
||||
source_model: Some("gpt-5.4-high"),
|
||||
family: SameFormatProviderFamily::Standard,
|
||||
|
||||
Reference in New Issue
Block a user