feat: add embedding and rerank support

This commit is contained in:
Kayphoon
2026-05-03 17:32:41 +08:00
parent 3e2eca4fd0
commit 5abe664d65
87 changed files with 5520 additions and 184 deletions

View File

@@ -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"
);
}
}
}

View File

@@ -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"
));
}
}

View File

@@ -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"
);
}
}

View File

@@ -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,