mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
feat: add embedding and rerank support
This commit is contained in:
@@ -3,6 +3,13 @@ use chrono::{SecondsFormat, Utc};
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
const EMBEDDING_API_FORMATS: &[&str] = &[
|
||||
"openai:embedding",
|
||||
"jina:embedding",
|
||||
"gemini:embedding",
|
||||
"doubao:embedding",
|
||||
];
|
||||
|
||||
fn unix_secs_to_rfc3339(unix_secs: u64) -> Option<String> {
|
||||
let timestamp = i64::try_from(unix_secs).ok()?;
|
||||
Some(
|
||||
@@ -36,6 +43,50 @@ fn model_effective_capability(
|
||||
})
|
||||
}
|
||||
|
||||
fn value_contains_string(value: &Value, expected: &str) -> bool {
|
||||
match value {
|
||||
Value::String(value) => value.trim().eq_ignore_ascii_case(expected),
|
||||
Value::Array(values) => values
|
||||
.iter()
|
||||
.any(|value| value_contains_string(value, expected)),
|
||||
Value::Object(object) => object
|
||||
.values()
|
||||
.any(|value| value_contains_string(value, expected)),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn value_has_true_key(value: &Value, key: &str) -> bool {
|
||||
value
|
||||
.as_object()
|
||||
.and_then(|object| object.get(key))
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn value_contains_embedding_metadata(value: &Value) -> bool {
|
||||
value_has_true_key(value, "embedding")
|
||||
|| value_contains_string(value, "embedding")
|
||||
|| EMBEDDING_API_FORMATS
|
||||
.iter()
|
||||
.any(|api_format| value_contains_string(value, api_format))
|
||||
}
|
||||
|
||||
fn model_effective_embedding_capability(model: &StoredAdminProviderModel) -> bool {
|
||||
model
|
||||
.config
|
||||
.as_ref()
|
||||
.is_some_and(value_contains_embedding_metadata)
|
||||
|| model
|
||||
.global_model_supported_capabilities
|
||||
.as_ref()
|
||||
.is_some_and(value_contains_embedding_metadata)
|
||||
|| model
|
||||
.global_model_config
|
||||
.as_ref()
|
||||
.is_some_and(value_contains_embedding_metadata)
|
||||
}
|
||||
|
||||
fn merge_json_values(base: &mut Value, overlay: Value) {
|
||||
match (base, overlay) {
|
||||
(Value::Object(base_map), Value::Object(overlay_map)) => {
|
||||
@@ -128,6 +179,7 @@ pub fn admin_provider_model_effective_capability(
|
||||
model.global_model_config.as_ref(),
|
||||
"image_generation",
|
||||
),
|
||||
"embedding" => model_effective_embedding_capability(model),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
@@ -161,6 +213,7 @@ pub fn build_admin_provider_model_response(
|
||||
"supports_streaming": model.supports_streaming,
|
||||
"supports_extended_thinking": model.supports_extended_thinking,
|
||||
"supports_image_generation": model.supports_image_generation,
|
||||
"supports_embedding": model_effective_embedding_capability(model),
|
||||
"effective_supports_vision": admin_provider_model_effective_capability(model, "vision"),
|
||||
"effective_supports_function_calling": admin_provider_model_effective_capability(
|
||||
model,
|
||||
@@ -175,6 +228,10 @@ pub fn build_admin_provider_model_response(
|
||||
model,
|
||||
"image_generation",
|
||||
),
|
||||
"effective_supports_embedding": admin_provider_model_effective_capability(
|
||||
model,
|
||||
"embedding",
|
||||
),
|
||||
"is_active": model.is_active,
|
||||
"is_available": model.is_available,
|
||||
"config": model.config.clone(),
|
||||
@@ -214,6 +271,7 @@ pub fn build_admin_provider_available_source_models_payload(
|
||||
"supports_vision": admin_provider_model_effective_capability(&model, "vision"),
|
||||
"supports_function_calling": admin_provider_model_effective_capability(&model, "function_calling"),
|
||||
"supports_streaming": admin_provider_model_effective_capability(&model, "streaming"),
|
||||
"supports_embedding": admin_provider_model_effective_capability(&model, "embedding"),
|
||||
}),
|
||||
"is_active": model.is_active,
|
||||
})
|
||||
|
||||
@@ -545,6 +545,18 @@ const ADMIN_API_FORMAT_DEFINITIONS: &[AdminApiFormatDefinition] = &[
|
||||
default_path: "/v1/responses/compact",
|
||||
aliases: &["responses_compact"],
|
||||
},
|
||||
AdminApiFormatDefinition {
|
||||
value: "openai:embedding",
|
||||
label: "OpenAI Embedding",
|
||||
default_path: "/v1/embeddings",
|
||||
aliases: &["openai_embedding", "embeddings"],
|
||||
},
|
||||
AdminApiFormatDefinition {
|
||||
value: "openai:rerank",
|
||||
label: "OpenAI Rerank",
|
||||
default_path: "/v1/rerank",
|
||||
aliases: &["openai_rerank", "rerank"],
|
||||
},
|
||||
AdminApiFormatDefinition {
|
||||
value: "openai:image",
|
||||
label: "OpenAI Image",
|
||||
@@ -569,12 +581,36 @@ const ADMIN_API_FORMAT_DEFINITIONS: &[AdminApiFormatDefinition] = &[
|
||||
default_path: "/v1beta/models/{model}:{action}",
|
||||
aliases: &["gemini", "google", "vertex"],
|
||||
},
|
||||
AdminApiFormatDefinition {
|
||||
value: "gemini:embedding",
|
||||
label: "Gemini Embedding",
|
||||
default_path: "/v1/embeddings",
|
||||
aliases: &["gemini_embedding"],
|
||||
},
|
||||
AdminApiFormatDefinition {
|
||||
value: "gemini:video",
|
||||
label: "Gemini Video",
|
||||
default_path: "/v1beta/models/{model}:predictLongRunning",
|
||||
aliases: &["gemini_video", "veo"],
|
||||
},
|
||||
AdminApiFormatDefinition {
|
||||
value: "jina:embedding",
|
||||
label: "Jina Embedding",
|
||||
default_path: "/v1/embeddings",
|
||||
aliases: &["jina_embedding"],
|
||||
},
|
||||
AdminApiFormatDefinition {
|
||||
value: "jina:rerank",
|
||||
label: "Jina Rerank",
|
||||
default_path: "/v1/rerank",
|
||||
aliases: &["jina_rerank"],
|
||||
},
|
||||
AdminApiFormatDefinition {
|
||||
value: "doubao:embedding",
|
||||
label: "Doubao Embedding",
|
||||
default_path: "/v1/embeddings",
|
||||
aliases: &["doubao_embedding"],
|
||||
},
|
||||
];
|
||||
|
||||
pub fn build_admin_system_check_update_payload(current_version: String) -> serde_json::Value {
|
||||
|
||||
@@ -23,9 +23,10 @@ pub use crate::contracts::{
|
||||
OPENAI_CHAT_STREAM_PLAN_KIND, OPENAI_CHAT_STREAM_SUCCESS_REPORT_KIND,
|
||||
OPENAI_CHAT_SYNC_ERROR_REPORT_KIND, OPENAI_CHAT_SYNC_FINALIZE_REPORT_KIND,
|
||||
OPENAI_CHAT_SYNC_PLAN_KIND, OPENAI_CHAT_SYNC_SUCCESS_REPORT_KIND,
|
||||
OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_STREAM_SUCCESS_REPORT_KIND,
|
||||
OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
|
||||
OPENAI_IMAGE_SYNC_SUCCESS_REPORT_KIND, OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND,
|
||||
OPENAI_EMBEDDING_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND,
|
||||
OPENAI_IMAGE_STREAM_SUCCESS_REPORT_KIND, OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND,
|
||||
OPENAI_IMAGE_SYNC_PLAN_KIND, OPENAI_IMAGE_SYNC_SUCCESS_REPORT_KIND,
|
||||
OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND,
|
||||
OPENAI_RESPONSES_COMPACT_STREAM_SUCCESS_REPORT_KIND,
|
||||
OPENAI_RESPONSES_COMPACT_SYNC_ERROR_REPORT_KIND,
|
||||
OPENAI_RESPONSES_COMPACT_SYNC_FINALIZE_REPORT_KIND, OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND,
|
||||
|
||||
@@ -18,12 +18,12 @@ pub use plan_kinds::{
|
||||
GEMINI_FILES_DOWNLOAD_PLAN_KIND, GEMINI_FILES_GET_PLAN_KIND, GEMINI_FILES_LIST_PLAN_KIND,
|
||||
GEMINI_FILES_UPLOAD_PLAN_KIND, GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND,
|
||||
GEMINI_VIDEO_CREATE_SYNC_PLAN_KIND, OPENAI_CHAT_STREAM_PLAN_KIND, OPENAI_CHAT_SYNC_PLAN_KIND,
|
||||
OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
|
||||
OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND, OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND,
|
||||
OPENAI_RESPONSES_STREAM_PLAN_KIND, OPENAI_RESPONSES_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND, OPENAI_VIDEO_CONTENT_PLAN_KIND,
|
||||
OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND, OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
|
||||
OPENAI_EMBEDDING_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
|
||||
OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND,
|
||||
OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND, OPENAI_RESPONSES_STREAM_PLAN_KIND,
|
||||
OPENAI_RESPONSES_SYNC_PLAN_KIND, OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_CONTENT_PLAN_KIND, OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND, OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
|
||||
};
|
||||
pub use report_kinds::{
|
||||
core_error_background_report_kind, core_error_default_client_api_format,
|
||||
|
||||
@@ -20,6 +20,8 @@ pub const CLAUDE_CLI_STREAM_PLAN_KIND: &str = "claude_cli_stream";
|
||||
pub const GEMINI_CLI_STREAM_PLAN_KIND: &str = "gemini_cli_stream";
|
||||
pub const OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND: &str = "openai_video_create_sync";
|
||||
pub const OPENAI_CHAT_SYNC_PLAN_KIND: &str = "openai_chat_sync";
|
||||
pub const OPENAI_EMBEDDING_SYNC_PLAN_KIND: &str = "openai_embedding_sync";
|
||||
pub const OPENAI_RERANK_SYNC_PLAN_KIND: &str = "openai_rerank_sync";
|
||||
pub const OPENAI_RESPONSES_SYNC_PLAN_KIND: &str = "openai_responses_sync";
|
||||
pub const OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND: &str = "openai_responses_compact_sync";
|
||||
pub const CLAUDE_CHAT_SYNC_PLAN_KIND: &str = "claude_chat_sync";
|
||||
|
||||
@@ -9,19 +9,21 @@ pub mod response;
|
||||
|
||||
pub use protocol::canonical::{
|
||||
canonical_request_unknown_block_count, canonical_response_unknown_block_count,
|
||||
canonical_to_claude_request, canonical_to_claude_response, canonical_to_gemini_request,
|
||||
canonical_to_gemini_response, canonical_to_openai_chat_request,
|
||||
canonical_to_claude_request, canonical_to_claude_response, canonical_to_embedding_response,
|
||||
canonical_to_gemini_request, canonical_to_gemini_response, canonical_to_openai_chat_request,
|
||||
canonical_to_openai_chat_response, canonical_to_openai_responses_compact_request,
|
||||
canonical_to_openai_responses_compact_response, canonical_to_openai_responses_request,
|
||||
canonical_to_openai_responses_response, canonical_unknown_block_count,
|
||||
from_claude_to_canonical_request, from_claude_to_canonical_response,
|
||||
from_gemini_to_canonical_request, from_gemini_to_canonical_response,
|
||||
from_openai_chat_to_canonical_request, from_openai_chat_to_canonical_response,
|
||||
from_openai_responses_to_canonical_request, from_openai_responses_to_canonical_response,
|
||||
CanonicalContentBlock, CanonicalGenerationConfig, CanonicalInstruction, CanonicalMessage,
|
||||
CanonicalRequest, CanonicalResponse, CanonicalResponseFormat, CanonicalResponseOutput,
|
||||
CanonicalRole, CanonicalStopReason, CanonicalStreamEvent, CanonicalStreamFrame,
|
||||
CanonicalThinkingConfig, CanonicalToolChoice, CanonicalToolDefinition, CanonicalUsage,
|
||||
from_embedding_to_canonical_response, from_gemini_to_canonical_request,
|
||||
from_gemini_to_canonical_response, from_openai_chat_to_canonical_request,
|
||||
from_openai_chat_to_canonical_response, from_openai_responses_to_canonical_request,
|
||||
from_openai_responses_to_canonical_response, CanonicalContentBlock, CanonicalEmbedding,
|
||||
CanonicalEmbeddingInput, CanonicalEmbeddingRequest, CanonicalEmbeddingResponse,
|
||||
CanonicalGenerationConfig, CanonicalInstruction, CanonicalMessage, CanonicalRequest,
|
||||
CanonicalResponse, CanonicalResponseFormat, CanonicalResponseOutput, CanonicalRole,
|
||||
CanonicalStopReason, CanonicalStreamEvent, CanonicalStreamFrame, CanonicalThinkingConfig,
|
||||
CanonicalToolChoice, CanonicalToolDefinition, CanonicalUsage,
|
||||
};
|
||||
pub use protocol::context::{FormatContext, FormatError};
|
||||
pub use protocol::formats::{
|
||||
@@ -30,10 +32,11 @@ pub use protocol::formats::{
|
||||
FormatFamily, FormatId, FormatProfile,
|
||||
};
|
||||
pub use protocol::matrix::{
|
||||
request_candidate_api_format_preference, request_candidate_api_formats,
|
||||
request_conversion_kind, request_conversion_requires_enable_flag,
|
||||
sync_chat_response_conversion_kind, sync_cli_response_conversion_kind, RequestConversionKind,
|
||||
SyncChatResponseConversionKind, SyncCliResponseConversionKind,
|
||||
is_embedding_api_format, is_rerank_api_format, request_candidate_api_format_preference,
|
||||
request_candidate_api_formats, request_conversion_kind,
|
||||
request_conversion_requires_enable_flag, sync_chat_response_conversion_kind,
|
||||
sync_cli_response_conversion_kind, RequestConversionKind, SyncChatResponseConversionKind,
|
||||
SyncCliResponseConversionKind,
|
||||
};
|
||||
pub use protocol::registry::{build_stream_transcoder, convert_request, convert_response};
|
||||
pub use request::model_directives::{
|
||||
|
||||
@@ -222,6 +222,94 @@ pub struct CanonicalUsage {
|
||||
pub extensions: BTreeMap<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum CanonicalEmbeddingInput {
|
||||
String(String),
|
||||
StringArray(Vec<String>),
|
||||
TokenArray(Vec<i64>),
|
||||
TokenArrayArray(Vec<Vec<i64>>),
|
||||
}
|
||||
|
||||
impl CanonicalEmbeddingInput {
|
||||
fn is_empty(&self) -> bool {
|
||||
match self {
|
||||
Self::String(value) => value.trim().is_empty(),
|
||||
Self::StringArray(values) => {
|
||||
values.is_empty() || values.iter().any(|value| value.trim().is_empty())
|
||||
}
|
||||
Self::TokenArray(values) => values.is_empty(),
|
||||
Self::TokenArrayArray(values) => values.is_empty() || values.iter().any(Vec::is_empty),
|
||||
}
|
||||
}
|
||||
|
||||
fn as_string_items(&self) -> Option<Vec<&str>> {
|
||||
match self {
|
||||
Self::String(value) => Some(vec![value.as_str()]),
|
||||
Self::StringArray(values) => Some(values.iter().map(String::as_str).collect()),
|
||||
Self::TokenArray(_) | Self::TokenArrayArray(_) => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct CanonicalEmbeddingRequest {
|
||||
pub input: CanonicalEmbeddingInput,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub encoding_format: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub dimensions: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub task: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub user: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
|
||||
pub extensions: BTreeMap<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct CanonicalRerankRequest {
|
||||
pub query: String,
|
||||
#[serde(default)]
|
||||
pub documents: Vec<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub top_n: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub return_documents: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
|
||||
pub extensions: BTreeMap<String, Value>,
|
||||
}
|
||||
|
||||
impl CanonicalRerankRequest {
|
||||
fn is_empty(&self) -> bool {
|
||||
self.query.trim().is_empty()
|
||||
|| self.documents.is_empty()
|
||||
|| self.documents.iter().any(rerank_document_is_empty)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct CanonicalEmbedding {
|
||||
#[serde(default)]
|
||||
pub index: usize,
|
||||
#[serde(default)]
|
||||
pub embedding: Vec<f64>,
|
||||
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
|
||||
pub extensions: BTreeMap<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct CanonicalEmbeddingResponse {
|
||||
pub id: String,
|
||||
pub model: String,
|
||||
#[serde(default)]
|
||||
pub embeddings: Vec<CanonicalEmbedding>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub usage: Option<CanonicalUsage>,
|
||||
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
|
||||
pub extensions: BTreeMap<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct CanonicalRequest {
|
||||
#[serde(default)]
|
||||
@@ -232,6 +320,10 @@ pub struct CanonicalRequest {
|
||||
pub system: Option<String>,
|
||||
#[serde(default)]
|
||||
pub messages: Vec<CanonicalMessage>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub embedding: Option<CanonicalEmbeddingRequest>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub rerank: Option<CanonicalRerankRequest>,
|
||||
#[serde(default)]
|
||||
pub generation: CanonicalGenerationConfig,
|
||||
#[serde(default)]
|
||||
@@ -373,6 +465,47 @@ pub fn canonical_to_gemini_request(
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn from_embedding_to_canonical_request(
|
||||
body_json: &Value,
|
||||
namespace: &str,
|
||||
) -> Option<CanonicalRequest> {
|
||||
embedding_request_from_raw(body_json, namespace)
|
||||
}
|
||||
|
||||
pub(crate) fn canonical_to_embedding_request(
|
||||
canonical: &CanonicalRequest,
|
||||
mapped_model: &str,
|
||||
namespace: &str,
|
||||
) -> Option<Value> {
|
||||
match namespace {
|
||||
"openai" => canonical_to_openai_embedding_request(canonical, mapped_model),
|
||||
"jina" => canonical_to_jina_embedding_request(canonical, mapped_model),
|
||||
"gemini" => canonical_to_gemini_embedding_request(canonical, mapped_model),
|
||||
"doubao" => canonical_to_doubao_embedding_request(canonical, mapped_model),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn from_rerank_to_canonical_request(
|
||||
body_json: &Value,
|
||||
namespace: &str,
|
||||
) -> Option<CanonicalRequest> {
|
||||
rerank_request_from_raw(body_json, namespace)
|
||||
}
|
||||
|
||||
pub(crate) fn canonical_to_rerank_request(
|
||||
canonical: &CanonicalRequest,
|
||||
mapped_model: &str,
|
||||
namespace: &str,
|
||||
) -> Option<Value> {
|
||||
match namespace {
|
||||
"openai" | "jina" => {
|
||||
canonical_to_openai_like_rerank_request(canonical, mapped_model, namespace)
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn from_openai_chat_to_canonical_response(body_json: &Value) -> Option<CanonicalResponse> {
|
||||
crate::protocol::formats::openai_chat::response::from_raw(body_json)
|
||||
}
|
||||
@@ -556,6 +689,23 @@ pub fn canonical_to_gemini_response(
|
||||
crate::protocol::formats::gemini_generate_content::response::to_raw(canonical, report_context)
|
||||
}
|
||||
|
||||
pub fn from_embedding_to_canonical_response(
|
||||
body_json: &Value,
|
||||
namespace: &str,
|
||||
) -> Option<CanonicalEmbeddingResponse> {
|
||||
embedding_response_from_raw(body_json, namespace)
|
||||
}
|
||||
|
||||
pub fn canonical_to_embedding_response(
|
||||
canonical: &CanonicalEmbeddingResponse,
|
||||
namespace: &str,
|
||||
) -> Option<Value> {
|
||||
match namespace {
|
||||
"openai" | "jina" => Some(canonical_to_openai_embedding_response(canonical, namespace)),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn canonical_unknown_block_count(blocks: &[CanonicalContentBlock]) -> usize {
|
||||
blocks
|
||||
.iter()
|
||||
@@ -4123,6 +4273,416 @@ pub(crate) fn strip_claude_billing_header(text: &str) -> String {
|
||||
remainder.trim_start_matches('\n').trim().to_string()
|
||||
}
|
||||
|
||||
fn embedding_request_from_raw(body_json: &Value, namespace: &str) -> Option<CanonicalRequest> {
|
||||
let request = body_json.as_object()?;
|
||||
let model = request
|
||||
.get("model")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())?
|
||||
.to_string();
|
||||
let input =
|
||||
serde_json::from_value::<CanonicalEmbeddingInput>(request.get("input")?.clone()).ok()?;
|
||||
if input.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let embedding = CanonicalEmbeddingRequest {
|
||||
input,
|
||||
encoding_format: request
|
||||
.get("encoding_format")
|
||||
.and_then(Value::as_str)
|
||||
.map(ToOwned::to_owned),
|
||||
dimensions: request.get("dimensions").and_then(Value::as_u64),
|
||||
task: request
|
||||
.get("task")
|
||||
.and_then(Value::as_str)
|
||||
.map(ToOwned::to_owned),
|
||||
user: request
|
||||
.get("user")
|
||||
.and_then(Value::as_str)
|
||||
.map(ToOwned::to_owned),
|
||||
extensions: namespace_extensions(
|
||||
namespace,
|
||||
request,
|
||||
&[
|
||||
"model",
|
||||
"input",
|
||||
"encoding_format",
|
||||
"dimensions",
|
||||
"task",
|
||||
"user",
|
||||
],
|
||||
),
|
||||
};
|
||||
|
||||
Some(CanonicalRequest {
|
||||
model,
|
||||
embedding: Some(embedding),
|
||||
..CanonicalRequest::default()
|
||||
})
|
||||
}
|
||||
|
||||
fn rerank_request_from_raw(body_json: &Value, namespace: &str) -> Option<CanonicalRequest> {
|
||||
let request = body_json.as_object()?;
|
||||
let model = request
|
||||
.get("model")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())?
|
||||
.to_string();
|
||||
let query = request
|
||||
.get("query")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())?
|
||||
.to_string();
|
||||
let documents = request
|
||||
.get("documents")
|
||||
.and_then(Value::as_array)?
|
||||
.iter()
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
let rerank = CanonicalRerankRequest {
|
||||
query,
|
||||
documents,
|
||||
top_n: request
|
||||
.get("top_n")
|
||||
.or_else(|| request.get("topN"))
|
||||
.and_then(Value::as_u64),
|
||||
return_documents: request
|
||||
.get("return_documents")
|
||||
.or_else(|| request.get("returnDocuments"))
|
||||
.and_then(Value::as_bool),
|
||||
extensions: namespace_extensions(
|
||||
namespace,
|
||||
request,
|
||||
&[
|
||||
"model",
|
||||
"query",
|
||||
"documents",
|
||||
"top_n",
|
||||
"topN",
|
||||
"return_documents",
|
||||
"returnDocuments",
|
||||
],
|
||||
),
|
||||
};
|
||||
if rerank.is_empty() || rerank.top_n == Some(0) {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(CanonicalRequest {
|
||||
model,
|
||||
rerank: Some(rerank),
|
||||
..CanonicalRequest::default()
|
||||
})
|
||||
}
|
||||
|
||||
fn canonical_to_openai_like_rerank_request(
|
||||
canonical: &CanonicalRequest,
|
||||
mapped_model: &str,
|
||||
namespace: &str,
|
||||
) -> Option<Value> {
|
||||
let rerank = canonical.rerank.as_ref()?;
|
||||
if rerank.is_empty() || rerank.top_n == Some(0) {
|
||||
return None;
|
||||
}
|
||||
let mut output = Map::new();
|
||||
output.insert(
|
||||
"model".to_string(),
|
||||
Value::String(mapped_rerank_model(canonical, mapped_model)),
|
||||
);
|
||||
output.insert("query".to_string(), Value::String(rerank.query.clone()));
|
||||
output.insert(
|
||||
"documents".to_string(),
|
||||
Value::Array(rerank.documents.clone()),
|
||||
);
|
||||
if let Some(value) = rerank.top_n {
|
||||
output.insert("top_n".to_string(), Value::from(value));
|
||||
}
|
||||
if let Some(value) = rerank.return_documents {
|
||||
output.insert("return_documents".to_string(), Value::Bool(value));
|
||||
}
|
||||
output.extend(namespace_extension_object(
|
||||
&rerank.extensions,
|
||||
namespace,
|
||||
&output,
|
||||
));
|
||||
Some(Value::Object(output))
|
||||
}
|
||||
|
||||
fn rerank_document_is_empty(value: &Value) -> bool {
|
||||
match value {
|
||||
Value::String(text) => text.trim().is_empty(),
|
||||
Value::Object(object) => object
|
||||
.get("text")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|text| text.trim().is_empty()),
|
||||
Value::Null => true,
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn mapped_rerank_model(canonical: &CanonicalRequest, mapped_model: &str) -> String {
|
||||
mapped_model
|
||||
.trim()
|
||||
.chars()
|
||||
.next()
|
||||
.map(|_| mapped_model.trim().to_string())
|
||||
.unwrap_or_else(|| canonical.model.clone())
|
||||
}
|
||||
|
||||
fn canonical_to_openai_embedding_request(
|
||||
canonical: &CanonicalRequest,
|
||||
mapped_model: &str,
|
||||
) -> Option<Value> {
|
||||
canonical_to_openai_like_embedding_request(canonical, mapped_model, "openai", false)
|
||||
}
|
||||
|
||||
fn canonical_to_jina_embedding_request(
|
||||
canonical: &CanonicalRequest,
|
||||
mapped_model: &str,
|
||||
) -> Option<Value> {
|
||||
canonical_to_openai_like_embedding_request(canonical, mapped_model, "jina", true)
|
||||
}
|
||||
|
||||
fn canonical_to_openai_like_embedding_request(
|
||||
canonical: &CanonicalRequest,
|
||||
mapped_model: &str,
|
||||
namespace: &str,
|
||||
default_task: bool,
|
||||
) -> Option<Value> {
|
||||
let embedding = canonical.embedding.as_ref()?;
|
||||
if embedding.input.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let mut output = Map::new();
|
||||
output.insert(
|
||||
"model".to_string(),
|
||||
Value::String(mapped_embedding_model(canonical, mapped_model)),
|
||||
);
|
||||
output.insert(
|
||||
"input".to_string(),
|
||||
serde_json::to_value(&embedding.input).ok()?,
|
||||
);
|
||||
if let Some(value) = &embedding.encoding_format {
|
||||
output.insert("encoding_format".to_string(), Value::String(value.clone()));
|
||||
}
|
||||
if let Some(value) = embedding.dimensions {
|
||||
output.insert("dimensions".to_string(), Value::from(value));
|
||||
}
|
||||
if let Some(value) = &embedding.user {
|
||||
output.insert("user".to_string(), Value::String(value.clone()));
|
||||
}
|
||||
if let Some(task) = embedding
|
||||
.task
|
||||
.as_ref()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
{
|
||||
output.insert("task".to_string(), Value::String(task.clone()));
|
||||
} else if default_task {
|
||||
output.insert(
|
||||
"task".to_string(),
|
||||
Value::String("text-matching".to_string()),
|
||||
);
|
||||
}
|
||||
output.extend(namespace_extension_object(
|
||||
&embedding.extensions,
|
||||
namespace,
|
||||
&output,
|
||||
));
|
||||
Some(Value::Object(output))
|
||||
}
|
||||
|
||||
fn canonical_to_gemini_embedding_request(
|
||||
canonical: &CanonicalRequest,
|
||||
mapped_model: &str,
|
||||
) -> Option<Value> {
|
||||
let embedding = canonical.embedding.as_ref()?;
|
||||
let items = embedding.input.as_string_items()?;
|
||||
if items.is_empty() || items.iter().any(|value| value.trim().is_empty()) {
|
||||
return None;
|
||||
}
|
||||
let model = mapped_embedding_model(canonical, mapped_model);
|
||||
if items.len() == 1 {
|
||||
return Some(json!({
|
||||
"model": model,
|
||||
"content": {
|
||||
"parts": [{"text": items[0]}]
|
||||
}
|
||||
}));
|
||||
}
|
||||
Some(json!({
|
||||
"model": model,
|
||||
"requests": items.into_iter().map(|text| {
|
||||
json!({
|
||||
"model": model,
|
||||
"content": {
|
||||
"parts": [{"text": text}]
|
||||
}
|
||||
})
|
||||
}).collect::<Vec<_>>()
|
||||
}))
|
||||
}
|
||||
|
||||
fn canonical_to_doubao_embedding_request(
|
||||
canonical: &CanonicalRequest,
|
||||
mapped_model: &str,
|
||||
) -> Option<Value> {
|
||||
let embedding = canonical.embedding.as_ref()?;
|
||||
let items = embedding.input.as_string_items()?;
|
||||
if items.is_empty() || items.iter().any(|value| value.trim().is_empty()) {
|
||||
return None;
|
||||
}
|
||||
let mut output = Map::new();
|
||||
output.insert(
|
||||
"model".to_string(),
|
||||
Value::String(mapped_embedding_model(canonical, mapped_model)),
|
||||
);
|
||||
output.insert(
|
||||
"input".to_string(),
|
||||
Value::Array(
|
||||
items
|
||||
.into_iter()
|
||||
.map(|text| json!({"type": "text", "text": text}))
|
||||
.collect(),
|
||||
),
|
||||
);
|
||||
if let Some(dimensions) = embedding.dimensions {
|
||||
output.insert("dimensions".to_string(), Value::from(dimensions));
|
||||
}
|
||||
output.extend(namespace_extension_object(
|
||||
&embedding.extensions,
|
||||
"doubao",
|
||||
&output,
|
||||
));
|
||||
Some(Value::Object(output))
|
||||
}
|
||||
|
||||
fn embedding_response_from_raw(
|
||||
body_json: &Value,
|
||||
namespace: &str,
|
||||
) -> Option<CanonicalEmbeddingResponse> {
|
||||
let body = body_json.as_object()?;
|
||||
if body.contains_key("error") {
|
||||
return None;
|
||||
}
|
||||
let data = body.get("data")?.as_array()?;
|
||||
let mut embeddings = Vec::new();
|
||||
for (fallback_index, item) in data.iter().enumerate() {
|
||||
let item_object = item.as_object()?;
|
||||
let values = item_object.get("embedding")?.as_array()?;
|
||||
let embedding = values
|
||||
.iter()
|
||||
.map(Value::as_f64)
|
||||
.collect::<Option<Vec<_>>>()?;
|
||||
embeddings.push(CanonicalEmbedding {
|
||||
index: item_object
|
||||
.get("index")
|
||||
.and_then(Value::as_u64)
|
||||
.and_then(|value| usize::try_from(value).ok())
|
||||
.unwrap_or(fallback_index),
|
||||
embedding,
|
||||
extensions: namespace_extensions(
|
||||
namespace,
|
||||
item_object,
|
||||
&["object", "index", "embedding"],
|
||||
),
|
||||
});
|
||||
}
|
||||
Some(CanonicalEmbeddingResponse {
|
||||
id: body
|
||||
.get("id")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("embd-unknown")
|
||||
.to_string(),
|
||||
model: body
|
||||
.get("model")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("unknown")
|
||||
.to_string(),
|
||||
embeddings,
|
||||
usage: openai_usage_to_canonical(body.get("usage")),
|
||||
extensions: namespace_extensions(
|
||||
namespace,
|
||||
body,
|
||||
&["id", "object", "model", "data", "usage"],
|
||||
),
|
||||
})
|
||||
}
|
||||
|
||||
fn canonical_to_openai_embedding_response(
|
||||
canonical: &CanonicalEmbeddingResponse,
|
||||
namespace: &str,
|
||||
) -> Value {
|
||||
let mut response = Map::new();
|
||||
response.insert("object".to_string(), Value::String("list".to_string()));
|
||||
if !canonical.model.trim().is_empty() && canonical.model != "unknown" {
|
||||
response.insert("model".to_string(), Value::String(canonical.model.clone()));
|
||||
}
|
||||
response.insert(
|
||||
"data".to_string(),
|
||||
Value::Array(
|
||||
canonical
|
||||
.embeddings
|
||||
.iter()
|
||||
.map(|embedding| {
|
||||
let mut item = Map::new();
|
||||
item.insert("object".to_string(), Value::String("embedding".to_string()));
|
||||
item.insert("index".to_string(), Value::from(embedding.index as u64));
|
||||
item.insert("embedding".to_string(), json!(embedding.embedding));
|
||||
item.extend(namespace_extension_object(
|
||||
&embedding.extensions,
|
||||
namespace,
|
||||
&item,
|
||||
));
|
||||
Value::Object(item)
|
||||
})
|
||||
.collect(),
|
||||
),
|
||||
);
|
||||
if let Some(usage) = &canonical.usage {
|
||||
response.insert("usage".to_string(), canonical_usage_to_openai(usage));
|
||||
}
|
||||
response.extend(namespace_extension_object(
|
||||
&canonical.extensions,
|
||||
namespace,
|
||||
&response,
|
||||
));
|
||||
Value::Object(response)
|
||||
}
|
||||
|
||||
fn mapped_embedding_model(canonical: &CanonicalRequest, mapped_model: &str) -> String {
|
||||
let mapped_model = mapped_model.trim();
|
||||
if mapped_model.is_empty() {
|
||||
canonical.model.clone()
|
||||
} else {
|
||||
mapped_model.to_string()
|
||||
}
|
||||
}
|
||||
|
||||
fn namespace_extensions(
|
||||
namespace: &str,
|
||||
object: &Map<String, Value>,
|
||||
handled_keys: &[&str],
|
||||
) -> BTreeMap<String, Value> {
|
||||
let handled = handled_keys
|
||||
.iter()
|
||||
.copied()
|
||||
.collect::<std::collections::BTreeSet<_>>();
|
||||
let raw = object
|
||||
.iter()
|
||||
.filter(|(key, _)| !handled.contains(key.as_str()))
|
||||
.map(|(key, value)| (key.clone(), value.clone()))
|
||||
.collect::<Map<String, Value>>();
|
||||
if raw.is_empty() {
|
||||
BTreeMap::new()
|
||||
} else {
|
||||
BTreeMap::from([(namespace.to_string(), Value::Object(raw))])
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
@@ -4135,10 +4695,328 @@ mod tests {
|
||||
from_gemini_to_canonical_request, from_gemini_to_canonical_response,
|
||||
from_openai_chat_to_canonical_request, from_openai_chat_to_canonical_response,
|
||||
from_openai_responses_to_canonical_request, from_openai_responses_to_canonical_response,
|
||||
CanonicalContentBlock, CanonicalRole,
|
||||
CanonicalContentBlock, CanonicalEmbedding, CanonicalEmbeddingInput,
|
||||
CanonicalEmbeddingRequest, CanonicalRole, CanonicalUsage,
|
||||
};
|
||||
use serde_json::{json, Value};
|
||||
|
||||
#[test]
|
||||
fn canonical_embedding_request_accepts_axonhub_input_shapes() {
|
||||
for input in [
|
||||
json!("hello"),
|
||||
json!(["hello", "world"]),
|
||||
json!([1, 2, 3]),
|
||||
json!([[1, 2], [3, 4]]),
|
||||
] {
|
||||
let request = json!({
|
||||
"model": "text-embedding-3-small",
|
||||
"embedding": {
|
||||
"input": input,
|
||||
"encoding_format": "float",
|
||||
"dimensions": 3
|
||||
}
|
||||
});
|
||||
let canonical = serde_json::from_value::<super::CanonicalRequest>(request)
|
||||
.expect("embedding request should deserialize");
|
||||
assert!(canonical.embedding.is_some());
|
||||
assert!(canonical.messages.is_empty());
|
||||
let encoded = serde_json::to_value(&canonical).expect("serialize");
|
||||
assert!(encoded.get("messages").is_some());
|
||||
assert!(encoded.get("embedding").is_some());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_wire_request_accepts_all_axonhub_input_shapes() {
|
||||
let cases = [
|
||||
(
|
||||
json!("hello"),
|
||||
"single string",
|
||||
CanonicalEmbeddingInput::String("hello".to_string()),
|
||||
),
|
||||
(
|
||||
json!(["hello", "world"]),
|
||||
"string array",
|
||||
CanonicalEmbeddingInput::StringArray(vec![
|
||||
"hello".to_string(),
|
||||
"world".to_string(),
|
||||
]),
|
||||
),
|
||||
(
|
||||
json!([1, 2, 3]),
|
||||
"token array",
|
||||
CanonicalEmbeddingInput::TokenArray(vec![1, 2, 3]),
|
||||
),
|
||||
(
|
||||
json!([[1, 2], [3, 4]]),
|
||||
"nested token array",
|
||||
CanonicalEmbeddingInput::TokenArrayArray(vec![vec![1, 2], vec![3, 4]]),
|
||||
),
|
||||
];
|
||||
|
||||
for (input, label, expected_input) in cases {
|
||||
let body = json!({
|
||||
"model": "text-embedding-3-small",
|
||||
"input": input
|
||||
});
|
||||
let canonical = super::from_embedding_to_canonical_request(&body, "openai")
|
||||
.unwrap_or_else(|| panic!("{label} should parse"));
|
||||
|
||||
assert_eq!(
|
||||
canonical.embedding.expect("embedding request").input,
|
||||
expected_input,
|
||||
"{label} should preserve its canonical variant"
|
||||
);
|
||||
assert!(canonical.messages.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_wire_request_rejects_empty_invalid_or_chat_payloads() {
|
||||
for body in [
|
||||
json!({"model": "text-embedding-3-small", "input": " "}),
|
||||
json!({"model": "text-embedding-3-small", "input": []}),
|
||||
json!({"model": "text-embedding-3-small", "input": [1, "two"]}),
|
||||
json!({"model": "text-embedding-3-small", "input": [[1], []]}),
|
||||
json!({"model": "", "input": "hello"}),
|
||||
json!({"input": "hello"}),
|
||||
json!({"model": "text-embedding-3-small", "messages": []}),
|
||||
] {
|
||||
assert!(
|
||||
super::from_embedding_to_canonical_request(&body, "openai").is_none(),
|
||||
"invalid embedding payload should be rejected: {body}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_openai_request_response_roundtrip_stays_non_chat() {
|
||||
let body = json!({
|
||||
"model": "text-embedding-3-small",
|
||||
"input": ["hello", "world"],
|
||||
"encoding_format": "float",
|
||||
"dimensions": 2,
|
||||
"user": "user-1",
|
||||
"extra": true
|
||||
});
|
||||
let canonical =
|
||||
super::from_embedding_to_canonical_request(&body, "openai").expect("embedding request");
|
||||
assert_eq!(canonical.model, "text-embedding-3-small");
|
||||
assert!(canonical.messages.is_empty());
|
||||
assert!(matches!(
|
||||
canonical.embedding.as_ref().map(|embedding| &embedding.input),
|
||||
Some(CanonicalEmbeddingInput::StringArray(values)) if values == &vec!["hello".to_string(), "world".to_string()]
|
||||
));
|
||||
|
||||
let rebuilt =
|
||||
super::canonical_to_embedding_request(&canonical, "upstream-embedding", "openai")
|
||||
.expect("openai embedding request");
|
||||
assert_eq!(rebuilt["model"], "upstream-embedding");
|
||||
assert_eq!(rebuilt["input"], json!(["hello", "world"]));
|
||||
assert!(rebuilt.get("messages").is_none());
|
||||
|
||||
let response = json!({
|
||||
"object": "list",
|
||||
"model": "upstream-embedding",
|
||||
"data": [
|
||||
{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]},
|
||||
{"object": "embedding", "index": 1, "embedding": [0.3, 0.4]}
|
||||
],
|
||||
"usage": {"prompt_tokens": 4, "total_tokens": 4}
|
||||
});
|
||||
let canonical_response = super::from_embedding_to_canonical_response(&response, "openai")
|
||||
.expect("embedding response");
|
||||
assert_eq!(canonical_response.embeddings.len(), 2);
|
||||
let emitted = super::canonical_to_embedding_response(&canonical_response, "openai")
|
||||
.expect("embedding response output");
|
||||
assert_eq!(emitted["data"][0]["embedding"], json!([0.1, 0.2]));
|
||||
assert!(emitted.get("choices").is_none());
|
||||
assert!(emitted.get("messages").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_provider_request_emitters_preserve_provider_contracts() {
|
||||
let canonical = super::CanonicalRequest {
|
||||
model: "text-embedding-3-small".to_string(),
|
||||
embedding: Some(CanonicalEmbeddingRequest {
|
||||
input: CanonicalEmbeddingInput::StringArray(vec![
|
||||
"alpha".to_string(),
|
||||
"beta".to_string(),
|
||||
]),
|
||||
encoding_format: Some("float".to_string()),
|
||||
dimensions: Some(2),
|
||||
task: None,
|
||||
user: None,
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let jina = super::canonical_to_embedding_request(&canonical, "jina-embeddings-v3", "jina")
|
||||
.expect("jina embedding request");
|
||||
assert_eq!(jina["task"], "text-matching");
|
||||
assert_eq!(jina["input"], json!(["alpha", "beta"]));
|
||||
|
||||
let gemini =
|
||||
super::canonical_to_embedding_request(&canonical, "gemini-embedding-001", "gemini")
|
||||
.expect("gemini embedding request");
|
||||
assert_eq!(
|
||||
gemini["requests"][0]["content"]["parts"][0]["text"],
|
||||
"alpha"
|
||||
);
|
||||
assert!(gemini.get("messages").is_none());
|
||||
|
||||
let doubao =
|
||||
super::canonical_to_embedding_request(&canonical, "doubao-embedding-vision", "doubao")
|
||||
.expect("doubao embedding request");
|
||||
assert_eq!(doubao["input"][0], json!({"type": "text", "text": "alpha"}));
|
||||
assert!(doubao.get("messages").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_provider_request_emitters_cover_golden_payload_variants() {
|
||||
let single = super::CanonicalRequest {
|
||||
model: "text-embedding-3-small".to_string(),
|
||||
embedding: Some(CanonicalEmbeddingRequest {
|
||||
input: CanonicalEmbeddingInput::String("alpha".to_string()),
|
||||
encoding_format: Some("float".to_string()),
|
||||
dimensions: Some(1536),
|
||||
task: Some("retrieval.passage".to_string()),
|
||||
user: Some("user-1".to_string()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let openai =
|
||||
super::canonical_to_embedding_request(&single, "text-embedding-3-large", "openai")
|
||||
.expect("openai embedding request");
|
||||
assert_eq!(openai["model"], "text-embedding-3-large");
|
||||
assert_eq!(openai["input"], "alpha");
|
||||
assert_eq!(openai["encoding_format"], "float");
|
||||
assert_eq!(openai["dimensions"], 1536);
|
||||
assert_eq!(openai["user"], "user-1");
|
||||
assert_eq!(openai["task"], "retrieval.passage");
|
||||
|
||||
let jina = super::canonical_to_embedding_request(&single, "jina-embeddings-v3", "jina")
|
||||
.expect("jina embedding request");
|
||||
assert_eq!(jina["task"], "retrieval.passage");
|
||||
assert_eq!(jina["input"], "alpha");
|
||||
|
||||
let gemini =
|
||||
super::canonical_to_embedding_request(&single, "gemini-embedding-001", "gemini")
|
||||
.expect("gemini single embedding request");
|
||||
assert_eq!(gemini["model"], "gemini-embedding-001");
|
||||
assert_eq!(gemini["content"]["parts"][0]["text"], "alpha");
|
||||
assert!(gemini.get("requests").is_none());
|
||||
|
||||
let doubao =
|
||||
super::canonical_to_embedding_request(&single, "doubao-embedding-vision", "doubao")
|
||||
.expect("doubao embedding request");
|
||||
assert_eq!(doubao["model"], "doubao-embedding-vision");
|
||||
assert_eq!(doubao["input"], json!([{"type": "text", "text": "alpha"}]));
|
||||
assert_eq!(doubao["dimensions"], 1536);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gemini_and_doubao_embedding_emitters_reject_token_inputs() {
|
||||
for input in [
|
||||
CanonicalEmbeddingInput::TokenArray(vec![1, 2, 3]),
|
||||
CanonicalEmbeddingInput::TokenArrayArray(vec![vec![1, 2], vec![3, 4]]),
|
||||
] {
|
||||
let canonical = super::CanonicalRequest {
|
||||
model: "token-model".to_string(),
|
||||
embedding: Some(CanonicalEmbeddingRequest {
|
||||
input,
|
||||
encoding_format: None,
|
||||
dimensions: None,
|
||||
task: None,
|
||||
user: None,
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
assert!(super::canonical_to_embedding_request(
|
||||
&canonical,
|
||||
"gemini-embedding-001",
|
||||
"gemini"
|
||||
)
|
||||
.is_none());
|
||||
assert!(super::canonical_to_embedding_request(
|
||||
&canonical,
|
||||
"doubao-embedding",
|
||||
"doubao"
|
||||
)
|
||||
.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_response_parser_rejects_error_and_malformed_vectors() {
|
||||
for body in [
|
||||
json!({"error": {"message": "bad"}}),
|
||||
json!({"object": "list"}),
|
||||
json!({"data": [{"object": "embedding", "embedding": [0.1, "bad"]}]}),
|
||||
json!({"data": [{"object": "embedding"}]}),
|
||||
] {
|
||||
assert!(
|
||||
super::from_embedding_to_canonical_response(&body, "openai").is_none(),
|
||||
"malformed embedding response should be rejected: {body}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_response_parser_uses_openai_fallback_fields() {
|
||||
let body = json!({
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"object": "embedding", "embedding": [0.1, 0.2]},
|
||||
{"object": "embedding", "index": 7, "embedding": [0.3, 0.4]}
|
||||
]
|
||||
});
|
||||
|
||||
let canonical = super::from_embedding_to_canonical_response(&body, "openai")
|
||||
.expect("fallback embedding response");
|
||||
assert_eq!(canonical.id, "embd-unknown");
|
||||
assert_eq!(canonical.model, "unknown");
|
||||
assert_eq!(canonical.embeddings[0].index, 0);
|
||||
assert_eq!(canonical.embeddings[1].index, 7);
|
||||
assert!(super::canonical_to_embedding_response(&canonical, "jina").is_some());
|
||||
assert!(super::canonical_to_embedding_response(&canonical, "gemini").is_none());
|
||||
assert!(super::canonical_to_embedding_response(&canonical, "doubao").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_response_contract_serializes_vectors_without_chat_outputs() {
|
||||
let response = super::CanonicalEmbeddingResponse {
|
||||
id: "embd-1".to_string(),
|
||||
model: "text-embedding-3-small".to_string(),
|
||||
embeddings: vec![CanonicalEmbedding {
|
||||
index: 0,
|
||||
embedding: vec![0.1, 0.2, 0.3],
|
||||
extensions: Default::default(),
|
||||
}],
|
||||
usage: Some(CanonicalUsage {
|
||||
input_tokens: 3,
|
||||
total_tokens: 3,
|
||||
..Default::default()
|
||||
}),
|
||||
extensions: Default::default(),
|
||||
};
|
||||
|
||||
let encoded = serde_json::to_value(&response).expect("serialize");
|
||||
assert_eq!(
|
||||
encoded["embeddings"][0]["embedding"],
|
||||
json!([0.1, 0.2, 0.3])
|
||||
);
|
||||
assert!(encoded.get("choices").is_none());
|
||||
let decoded = serde_json::from_value::<super::CanonicalEmbeddingResponse>(encoded)
|
||||
.expect("deserialize");
|
||||
assert_eq!(decoded, response);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn canonical_request_preserves_openai_multimodal_tools_and_extensions() {
|
||||
let request = json!({
|
||||
|
||||
@@ -16,6 +16,8 @@ pub enum FormatFamily {
|
||||
OpenAi,
|
||||
Claude,
|
||||
Gemini,
|
||||
Jina,
|
||||
Doubao,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
|
||||
@@ -29,8 +31,14 @@ pub enum FormatId {
|
||||
OpenAiChat,
|
||||
OpenAiResponses,
|
||||
OpenAiResponsesCompact,
|
||||
OpenAiEmbedding,
|
||||
OpenAiRerank,
|
||||
ClaudeMessages,
|
||||
GeminiGenerateContent,
|
||||
GeminiEmbedding,
|
||||
JinaEmbedding,
|
||||
JinaRerank,
|
||||
DoubaoEmbedding,
|
||||
}
|
||||
|
||||
impl FormatId {
|
||||
@@ -44,11 +52,15 @@ impl FormatId {
|
||||
|
||||
pub fn family(self) -> FormatFamily {
|
||||
match self {
|
||||
Self::OpenAiChat | Self::OpenAiResponses | Self::OpenAiResponsesCompact => {
|
||||
FormatFamily::OpenAi
|
||||
}
|
||||
Self::OpenAiChat
|
||||
| Self::OpenAiResponses
|
||||
| Self::OpenAiResponsesCompact
|
||||
| Self::OpenAiEmbedding
|
||||
| Self::OpenAiRerank => FormatFamily::OpenAi,
|
||||
Self::ClaudeMessages => FormatFamily::Claude,
|
||||
Self::GeminiGenerateContent => FormatFamily::Gemini,
|
||||
Self::GeminiGenerateContent | Self::GeminiEmbedding => FormatFamily::Gemini,
|
||||
Self::JinaEmbedding | Self::JinaRerank => FormatFamily::Jina,
|
||||
Self::DoubaoEmbedding => FormatFamily::Doubao,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -64,8 +76,14 @@ impl FormatId {
|
||||
Self::OpenAiChat => "openai:chat",
|
||||
Self::OpenAiResponses => "openai:responses",
|
||||
Self::OpenAiResponsesCompact => "openai:responses:compact",
|
||||
Self::OpenAiEmbedding => "openai:embedding",
|
||||
Self::OpenAiRerank => "openai:rerank",
|
||||
Self::ClaudeMessages => "claude:messages",
|
||||
Self::GeminiGenerateContent => "gemini:generate_content",
|
||||
Self::GeminiEmbedding => "gemini:embedding",
|
||||
Self::JinaEmbedding => "jina:embedding",
|
||||
Self::JinaRerank => "jina:rerank",
|
||||
Self::DoubaoEmbedding => "doubao:embedding",
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -86,8 +104,14 @@ impl FromStr for FormatId {
|
||||
"openai:responses:compact" | "/v1/responses/compact" => {
|
||||
Ok(Self::OpenAiResponsesCompact)
|
||||
}
|
||||
"openai:embedding" | "/v1/embeddings" => Ok(Self::OpenAiEmbedding),
|
||||
"openai:rerank" | "/v1/rerank" => Ok(Self::OpenAiRerank),
|
||||
"claude:messages" | "/v1/messages" => Ok(Self::ClaudeMessages),
|
||||
"gemini:generate_content" => Ok(Self::GeminiGenerateContent),
|
||||
"gemini:embedding" => Ok(Self::GeminiEmbedding),
|
||||
"jina:embedding" | "/jina/v1/embeddings" => Ok(Self::JinaEmbedding),
|
||||
"jina:rerank" | "/jina/v1/rerank" => Ok(Self::JinaRerank),
|
||||
"doubao:embedding" => Ok(Self::DoubaoEmbedding),
|
||||
_ => Err(()),
|
||||
}
|
||||
}
|
||||
@@ -136,6 +160,90 @@ mod tests {
|
||||
assert_eq!(FormatId::parse("gemini:cli"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_embedding_api_formats() {
|
||||
assert_eq!(
|
||||
FormatId::parse("openai:embedding"),
|
||||
Some(FormatId::OpenAiEmbedding)
|
||||
);
|
||||
assert_eq!(
|
||||
FormatId::parse("/v1/embeddings"),
|
||||
Some(FormatId::OpenAiEmbedding)
|
||||
);
|
||||
assert_eq!(
|
||||
FormatId::parse("gemini:embedding"),
|
||||
Some(FormatId::GeminiEmbedding)
|
||||
);
|
||||
assert_eq!(
|
||||
FormatId::parse("jina:embedding"),
|
||||
Some(FormatId::JinaEmbedding)
|
||||
);
|
||||
assert_eq!(
|
||||
FormatId::parse("/jina/v1/embeddings"),
|
||||
Some(FormatId::JinaEmbedding)
|
||||
);
|
||||
assert_eq!(
|
||||
FormatId::parse("doubao:embedding"),
|
||||
Some(FormatId::DoubaoEmbedding)
|
||||
);
|
||||
assert_eq!(FormatId::OpenAiEmbedding.to_string(), "openai:embedding");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_format_ids_keep_provider_family_and_default_profile() {
|
||||
use super::{FormatFamily, FormatProfile};
|
||||
|
||||
for (format, family) in [
|
||||
(FormatId::OpenAiEmbedding, FormatFamily::OpenAi),
|
||||
(FormatId::GeminiEmbedding, FormatFamily::Gemini),
|
||||
(FormatId::JinaEmbedding, FormatFamily::Jina),
|
||||
(FormatId::DoubaoEmbedding, FormatFamily::Doubao),
|
||||
] {
|
||||
assert_eq!(format.family(), family);
|
||||
assert_eq!(format.profile(), FormatProfile::Default);
|
||||
assert_eq!(FormatId::parse(format.as_str()), Some(format));
|
||||
assert_eq!(format.to_string(), format.as_str());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_rerank_api_formats() {
|
||||
assert_eq!(
|
||||
FormatId::parse("openai:rerank"),
|
||||
Some(FormatId::OpenAiRerank)
|
||||
);
|
||||
assert_eq!(FormatId::parse("/v1/rerank"), Some(FormatId::OpenAiRerank));
|
||||
assert_eq!(FormatId::parse("jina:rerank"), Some(FormatId::JinaRerank));
|
||||
assert_eq!(
|
||||
FormatId::parse("/jina/v1/rerank"),
|
||||
Some(FormatId::JinaRerank)
|
||||
);
|
||||
assert_eq!(FormatId::OpenAiRerank.to_string(), "openai:rerank");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rerank_format_ids_keep_provider_family_and_default_profile() {
|
||||
use super::{FormatFamily, FormatProfile};
|
||||
|
||||
for (format, family) in [
|
||||
(FormatId::OpenAiRerank, FormatFamily::OpenAi),
|
||||
(FormatId::JinaRerank, FormatFamily::Jina),
|
||||
] {
|
||||
assert_eq!(format.family(), family);
|
||||
assert_eq!(format.profile(), FormatProfile::Default);
|
||||
assert_eq!(FormatId::parse(format.as_str()), Some(format));
|
||||
assert_eq!(format.to_string(), format.as_str());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_unknown_embedding_format() {
|
||||
assert_eq!(FormatId::parse("embedding"), None);
|
||||
assert_eq!(FormatId::parse("openai:embeddings"), None);
|
||||
assert_eq!(FormatId::parse("claude:embedding"), None);
|
||||
assert_eq!(FormatId::parse("gemini:embed_content"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalizes_api_format_aliases() {
|
||||
assert_eq!(
|
||||
@@ -154,6 +262,10 @@ mod tests {
|
||||
normalize_api_format_alias("GEMINI:GENERATE_CONTENT"),
|
||||
"gemini:generate_content"
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_api_format_alias("OPENAI:EMBEDDING"),
|
||||
"openai:embedding"
|
||||
);
|
||||
assert_eq!(normalize_api_format_alias("openai:image"), "openai:image");
|
||||
assert_eq!(normalize_api_format_alias("openai:video"), "openai:video");
|
||||
assert_eq!(normalize_api_format_alias("gemini:video"), "gemini:video");
|
||||
@@ -184,5 +296,21 @@ mod tests {
|
||||
api_format_storage_aliases("gemini:generate_content"),
|
||||
vec!["gemini:generate_content".to_string()]
|
||||
);
|
||||
assert_eq!(
|
||||
api_format_storage_aliases("openai:embedding"),
|
||||
vec!["openai:embedding".to_string()]
|
||||
);
|
||||
assert_eq!(
|
||||
api_format_storage_aliases("gemini:embedding"),
|
||||
vec!["gemini:embedding".to_string()]
|
||||
);
|
||||
assert_eq!(
|
||||
api_format_storage_aliases("jina:embedding"),
|
||||
vec!["jina:embedding".to_string()]
|
||||
);
|
||||
assert_eq!(
|
||||
api_format_storage_aliases("doubao:embedding"),
|
||||
vec!["doubao:embedding".to_string()]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -37,6 +37,13 @@ const STANDARD_API_FORMAT_ORDER: &[&str] = &[
|
||||
"claude:messages",
|
||||
"gemini:generate_content",
|
||||
];
|
||||
const EMBEDDING_CANDIDATE_API_FORMATS: &[&str] = &[
|
||||
"openai:embedding",
|
||||
"gemini:embedding",
|
||||
"jina:embedding",
|
||||
"doubao:embedding",
|
||||
];
|
||||
const RERANK_CANDIDATE_API_FORMATS: &[&str] = &["openai:rerank", "jina:rerank"];
|
||||
|
||||
pub fn request_candidate_api_format_preference(
|
||||
client_api_format: &str,
|
||||
@@ -48,6 +55,26 @@ pub fn request_candidate_api_format_preference(
|
||||
if client_api_format == "openai:responses:compact" {
|
||||
return (provider_api_format == "openai:responses:compact").then_some((0, 0));
|
||||
}
|
||||
if is_embedding_api_format(client_api_format.as_str()) {
|
||||
return is_embedding_api_format(provider_api_format.as_str()).then_some((
|
||||
if client_api_format == provider_api_format {
|
||||
0
|
||||
} else {
|
||||
1
|
||||
},
|
||||
embedding_api_format_priority(provider_api_format.as_str()),
|
||||
));
|
||||
}
|
||||
if is_rerank_api_format(client_api_format.as_str()) {
|
||||
return is_rerank_api_format(provider_api_format.as_str()).then_some((
|
||||
if client_api_format == provider_api_format {
|
||||
0
|
||||
} else {
|
||||
1
|
||||
},
|
||||
rerank_api_format_priority(provider_api_format.as_str()),
|
||||
));
|
||||
}
|
||||
|
||||
let (client_family, client_kind) =
|
||||
parse_non_compact_standard_api_format(client_api_format.as_str())?;
|
||||
@@ -77,6 +104,22 @@ pub fn request_candidate_api_formats(
|
||||
if client_api_format == "openai:responses:compact" {
|
||||
return vec!["openai:responses:compact"];
|
||||
}
|
||||
if is_embedding_api_format(client_api_format.as_str()) {
|
||||
let mut candidate_api_formats = EMBEDDING_CANDIDATE_API_FORMATS.to_vec();
|
||||
candidate_api_formats.sort_by_key(|provider_api_format| {
|
||||
request_candidate_api_format_preference(client_api_format.as_str(), provider_api_format)
|
||||
.unwrap_or((u8::MAX, u8::MAX))
|
||||
});
|
||||
return candidate_api_formats;
|
||||
}
|
||||
if is_rerank_api_format(client_api_format.as_str()) {
|
||||
let mut candidate_api_formats = RERANK_CANDIDATE_API_FORMATS.to_vec();
|
||||
candidate_api_formats.sort_by_key(|provider_api_format| {
|
||||
request_candidate_api_format_preference(client_api_format.as_str(), provider_api_format)
|
||||
.unwrap_or((u8::MAX, u8::MAX))
|
||||
});
|
||||
return candidate_api_formats;
|
||||
}
|
||||
if parse_non_compact_standard_api_format(client_api_format.as_str()).is_none() {
|
||||
return Vec::new();
|
||||
}
|
||||
@@ -192,6 +235,20 @@ pub fn is_standard_api_format(api_format: &str) -> bool {
|
||||
)
|
||||
}
|
||||
|
||||
pub fn is_embedding_api_format(api_format: &str) -> bool {
|
||||
matches!(
|
||||
normalize_api_format_alias(api_format).as_str(),
|
||||
"openai:embedding" | "gemini:embedding" | "jina:embedding" | "doubao:embedding"
|
||||
)
|
||||
}
|
||||
|
||||
pub fn is_rerank_api_format(api_format: &str) -> bool {
|
||||
matches!(
|
||||
normalize_api_format_alias(api_format).as_str(),
|
||||
"openai:rerank" | "jina:rerank"
|
||||
)
|
||||
}
|
||||
|
||||
pub fn parse_non_compact_standard_api_format(
|
||||
api_format: &str,
|
||||
) -> Option<(&'static str, &'static str)> {
|
||||
@@ -210,6 +267,10 @@ pub fn api_data_format_id(api_format: &str) -> Option<&'static str> {
|
||||
"gemini:generate_content" => Some("gemini"),
|
||||
"openai:chat" => Some("openai_chat"),
|
||||
"openai:responses" | "openai:responses:compact" => Some("openai_responses"),
|
||||
"openai:embedding" | "gemini:embedding" | "jina:embedding" | "doubao:embedding" => {
|
||||
Some("embedding")
|
||||
}
|
||||
"openai:rerank" | "jina:rerank" => Some("rerank"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
@@ -226,9 +287,26 @@ fn standard_api_format_priority(api_format: &str) -> u8 {
|
||||
.unwrap_or(STANDARD_API_FORMAT_ORDER.len()) as u8
|
||||
}
|
||||
|
||||
fn embedding_api_format_priority(api_format: &str) -> u8 {
|
||||
let api_format = normalize_api_format_alias(api_format);
|
||||
EMBEDDING_CANDIDATE_API_FORMATS
|
||||
.iter()
|
||||
.position(|candidate| *candidate == api_format)
|
||||
.unwrap_or(EMBEDDING_CANDIDATE_API_FORMATS.len()) as u8
|
||||
}
|
||||
|
||||
fn rerank_api_format_priority(api_format: &str) -> u8 {
|
||||
let api_format = normalize_api_format_alias(api_format);
|
||||
RERANK_CANDIDATE_API_FORMATS
|
||||
.iter()
|
||||
.position(|candidate| *candidate == api_format)
|
||||
.unwrap_or(RERANK_CANDIDATE_API_FORMATS.len()) as u8
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
api_data_format_id, is_embedding_api_format, is_rerank_api_format,
|
||||
request_candidate_api_format_preference, request_candidate_api_formats,
|
||||
request_conversion_kind, request_conversion_requires_enable_flag,
|
||||
sync_chat_response_conversion_kind, sync_cli_response_conversion_kind,
|
||||
@@ -355,6 +433,162 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_candidate_registry_excludes_chat_generation_formats() {
|
||||
assert_eq!(
|
||||
request_candidate_api_formats("openai:embedding", false),
|
||||
vec![
|
||||
"openai:embedding",
|
||||
"gemini:embedding",
|
||||
"jina:embedding",
|
||||
"doubao:embedding",
|
||||
]
|
||||
);
|
||||
assert_eq!(
|
||||
request_candidate_api_formats("jina:embedding", false),
|
||||
vec![
|
||||
"jina:embedding",
|
||||
"openai:embedding",
|
||||
"gemini:embedding",
|
||||
"doubao:embedding",
|
||||
]
|
||||
);
|
||||
assert!(!request_candidate_api_formats("openai:embedding", false).contains(&"openai:chat"));
|
||||
assert!(!request_candidate_api_formats("openai:embedding", false)
|
||||
.contains(&"gemini:generate_content"));
|
||||
assert_eq!(
|
||||
request_conversion_kind("openai:embedding", "jina:embedding"),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
request_conversion_kind("openai:embedding", "openai:chat"),
|
||||
None
|
||||
);
|
||||
assert!(!request_conversion_requires_enable_flag(
|
||||
"openai:embedding",
|
||||
"jina:embedding"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_candidate_registry_covers_all_provider_orderings() {
|
||||
assert_eq!(
|
||||
request_candidate_api_formats("gemini:embedding", true),
|
||||
vec![
|
||||
"gemini:embedding",
|
||||
"openai:embedding",
|
||||
"jina:embedding",
|
||||
"doubao:embedding",
|
||||
]
|
||||
);
|
||||
assert_eq!(
|
||||
request_candidate_api_formats("doubao:embedding", false),
|
||||
vec![
|
||||
"doubao:embedding",
|
||||
"openai:embedding",
|
||||
"gemini:embedding",
|
||||
"jina:embedding",
|
||||
]
|
||||
);
|
||||
|
||||
let embedding_formats = [
|
||||
"openai:embedding",
|
||||
"gemini:embedding",
|
||||
"jina:embedding",
|
||||
"doubao:embedding",
|
||||
];
|
||||
for client_api_format in embedding_formats {
|
||||
for provider_api_format in embedding_formats {
|
||||
assert!(
|
||||
request_candidate_api_format_preference(client_api_format, provider_api_format)
|
||||
.is_some(),
|
||||
"{client_api_format} should consider {provider_api_format} as embedding candidate"
|
||||
);
|
||||
assert_eq!(
|
||||
request_conversion_kind(client_api_format, provider_api_format),
|
||||
None,
|
||||
"embedding pair should not use chat/generation conversion kind"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_candidate_registry_never_crosses_chat_generation_boundary() {
|
||||
let embedding_formats = [
|
||||
"openai:embedding",
|
||||
"gemini:embedding",
|
||||
"jina:embedding",
|
||||
"doubao:embedding",
|
||||
];
|
||||
let standard_formats = [
|
||||
"openai:chat",
|
||||
"openai:responses",
|
||||
"claude:messages",
|
||||
"gemini:generate_content",
|
||||
];
|
||||
|
||||
for embedding_api_format in embedding_formats {
|
||||
assert!(is_embedding_api_format(embedding_api_format));
|
||||
assert_eq!(api_data_format_id(embedding_api_format), Some("embedding"));
|
||||
for standard_api_format in standard_formats {
|
||||
assert_eq!(
|
||||
request_candidate_api_format_preference(
|
||||
embedding_api_format,
|
||||
standard_api_format
|
||||
),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
request_candidate_api_format_preference(
|
||||
standard_api_format,
|
||||
embedding_api_format
|
||||
),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
request_conversion_kind(embedding_api_format, standard_api_format),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
request_conversion_kind(standard_api_format, embedding_api_format),
|
||||
None
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rerank_candidate_registry_excludes_chat_and_embedding_formats() {
|
||||
assert_eq!(
|
||||
request_candidate_api_formats("openai:rerank", false),
|
||||
vec!["openai:rerank", "jina:rerank"]
|
||||
);
|
||||
assert_eq!(
|
||||
request_candidate_api_formats("jina:rerank", false),
|
||||
vec!["jina:rerank", "openai:rerank"]
|
||||
);
|
||||
assert_eq!(api_data_format_id("openai:rerank"), Some("rerank"));
|
||||
assert!(is_rerank_api_format("jina:rerank"));
|
||||
assert!(!is_embedding_api_format("openai:rerank"));
|
||||
assert_eq!(
|
||||
request_candidate_api_format_preference("openai:rerank", "openai:embedding"),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
request_candidate_api_format_preference("openai:rerank", "openai:chat"),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
request_conversion_kind("openai:rerank", "jina:rerank"),
|
||||
None
|
||||
);
|
||||
assert!(!request_conversion_requires_enable_flag(
|
||||
"openai:rerank",
|
||||
"jina:rerank"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_candidate_registry_prefers_same_kind_before_same_family_fallbacks() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -1,7 +1,11 @@
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
protocol::canonical::{CanonicalRequest, CanonicalResponse},
|
||||
protocol::canonical::{
|
||||
canonical_to_embedding_request, canonical_to_rerank_request,
|
||||
from_embedding_to_canonical_request, from_rerank_to_canonical_request, CanonicalRequest,
|
||||
CanonicalResponse,
|
||||
},
|
||||
protocol::formats::{
|
||||
claude_messages, gemini_generate_content, openai_chat, openai_responses, FormatId,
|
||||
},
|
||||
@@ -22,6 +26,11 @@ pub fn parse_request(
|
||||
}
|
||||
FormatId::ClaudeMessages => claude_messages::request::from(body, ctx),
|
||||
FormatId::GeminiGenerateContent => gemini_generate_content::request::from(body, ctx),
|
||||
FormatId::OpenAiEmbedding => from_embedding_to_canonical_request(body, "openai"),
|
||||
FormatId::JinaEmbedding => from_embedding_to_canonical_request(body, "jina"),
|
||||
FormatId::OpenAiRerank => from_rerank_to_canonical_request(body, "openai"),
|
||||
FormatId::JinaRerank => from_rerank_to_canonical_request(body, "jina"),
|
||||
FormatId::GeminiEmbedding | FormatId::DoubaoEmbedding => None,
|
||||
}
|
||||
.ok_or_else(|| FormatError::RequestParseFailed {
|
||||
format: source.as_str().to_string(),
|
||||
@@ -48,6 +57,36 @@ pub fn emit_request(
|
||||
FormatId::OpenAiResponsesCompact => openai_responses::request::to_compact(&request, ctx),
|
||||
FormatId::ClaudeMessages => claude_messages::request::to(&request, ctx),
|
||||
FormatId::GeminiGenerateContent => gemini_generate_content::request::to(&request, ctx),
|
||||
FormatId::OpenAiEmbedding => canonical_to_embedding_request(
|
||||
&request,
|
||||
ctx.mapped_model_or(request.model.as_str()),
|
||||
"openai",
|
||||
),
|
||||
FormatId::JinaEmbedding => canonical_to_embedding_request(
|
||||
&request,
|
||||
ctx.mapped_model_or(request.model.as_str()),
|
||||
"jina",
|
||||
),
|
||||
FormatId::OpenAiRerank => canonical_to_rerank_request(
|
||||
&request,
|
||||
ctx.mapped_model_or(request.model.as_str()),
|
||||
"openai",
|
||||
),
|
||||
FormatId::JinaRerank => canonical_to_rerank_request(
|
||||
&request,
|
||||
ctx.mapped_model_or(request.model.as_str()),
|
||||
"jina",
|
||||
),
|
||||
FormatId::GeminiEmbedding => canonical_to_embedding_request(
|
||||
&request,
|
||||
ctx.mapped_model_or(request.model.as_str()),
|
||||
"gemini",
|
||||
),
|
||||
FormatId::DoubaoEmbedding => canonical_to_embedding_request(
|
||||
&request,
|
||||
ctx.mapped_model_or(request.model.as_str()),
|
||||
"doubao",
|
||||
),
|
||||
}
|
||||
.ok_or_else(|| FormatError::RequestEmitFailed {
|
||||
format: target.as_str().to_string(),
|
||||
@@ -77,6 +116,12 @@ pub fn parse_response(
|
||||
}
|
||||
FormatId::ClaudeMessages => claude_messages::response::from(body, ctx),
|
||||
FormatId::GeminiGenerateContent => gemini_generate_content::response::from(body, ctx),
|
||||
FormatId::OpenAiEmbedding
|
||||
| FormatId::JinaEmbedding
|
||||
| FormatId::OpenAiRerank
|
||||
| FormatId::JinaRerank
|
||||
| FormatId::GeminiEmbedding
|
||||
| FormatId::DoubaoEmbedding => None,
|
||||
}
|
||||
.ok_or_else(|| FormatError::ResponseParseFailed {
|
||||
format: source.as_str().to_string(),
|
||||
@@ -95,6 +140,12 @@ pub fn emit_response(
|
||||
FormatId::OpenAiResponsesCompact => openai_responses::response::to_compact(response, ctx),
|
||||
FormatId::ClaudeMessages => claude_messages::response::to(response, ctx),
|
||||
FormatId::GeminiGenerateContent => gemini_generate_content::response::to(response, ctx),
|
||||
FormatId::OpenAiEmbedding
|
||||
| FormatId::JinaEmbedding
|
||||
| FormatId::OpenAiRerank
|
||||
| FormatId::JinaRerank
|
||||
| FormatId::GeminiEmbedding
|
||||
| FormatId::DoubaoEmbedding => None,
|
||||
}
|
||||
.ok_or_else(|| FormatError::ResponseEmitFailed {
|
||||
format: target.as_str().to_string(),
|
||||
@@ -169,6 +220,116 @@ mod tests {
|
||||
assert_eq!(converted["input"][0]["content"][0]["type"], "input_text");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn converts_openai_embedding_to_jina_without_chat_fields() {
|
||||
let body = json!({
|
||||
"model": "text-embedding-3-small",
|
||||
"input": ["alpha", "beta"],
|
||||
"dimensions": 2
|
||||
});
|
||||
let ctx = FormatContext::default().with_mapped_model("jina-embeddings-v3");
|
||||
|
||||
let converted = convert_request("openai:embedding", "jina:embedding", &body, &ctx)
|
||||
.expect("embedding request conversion should succeed");
|
||||
|
||||
assert_eq!(converted["model"], "jina-embeddings-v3");
|
||||
assert_eq!(converted["task"], "text-matching");
|
||||
assert_eq!(converted["input"], json!(["alpha", "beta"]));
|
||||
assert!(converted.get("messages").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn converts_openai_embedding_to_gemini_and_doubao_payload_shapes() {
|
||||
let body = json!({
|
||||
"model": "text-embedding-3-small",
|
||||
"input": ["alpha", "beta"],
|
||||
"dimensions": 2
|
||||
});
|
||||
|
||||
let gemini = convert_request(
|
||||
"openai:embedding",
|
||||
"gemini:embedding",
|
||||
&body,
|
||||
&FormatContext::default().with_mapped_model("gemini-embedding-001"),
|
||||
)
|
||||
.expect("gemini embedding conversion should succeed");
|
||||
assert_eq!(gemini["model"], "gemini-embedding-001");
|
||||
assert_eq!(
|
||||
gemini["requests"][0]["content"]["parts"][0]["text"],
|
||||
"alpha"
|
||||
);
|
||||
assert!(gemini.get("messages").is_none());
|
||||
|
||||
let doubao = convert_request(
|
||||
"openai:embedding",
|
||||
"doubao:embedding",
|
||||
&body,
|
||||
&FormatContext::default().with_mapped_model("doubao-embedding-vision"),
|
||||
)
|
||||
.expect("doubao embedding conversion should succeed");
|
||||
assert_eq!(doubao["model"], "doubao-embedding-vision");
|
||||
assert_eq!(doubao["input"][0], json!({"type": "text", "text": "alpha"}));
|
||||
assert!(doubao.get("messages").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_registry_keeps_gemini_and_doubao_emit_only() {
|
||||
let body = json!({
|
||||
"model": "gemini-embedding-001",
|
||||
"content": {"parts": [{"text": "alpha"}]}
|
||||
});
|
||||
let ctx = FormatContext::default();
|
||||
|
||||
assert!(convert_request("gemini:embedding", "openai:embedding", &body, &ctx).is_err());
|
||||
assert!(convert_request("doubao:embedding", "openai:embedding", &body, &ctx).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_registry_rejects_chat_payload_for_embedding_format() {
|
||||
let body = json!({
|
||||
"model": "gpt-5",
|
||||
"messages": [{"role": "user", "content": "hello"}]
|
||||
});
|
||||
let ctx = FormatContext::default();
|
||||
|
||||
assert!(convert_request("openai:embedding", "jina:embedding", &body, &ctx).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn converts_openai_rerank_to_jina_without_chat_fields() {
|
||||
let body = json!({
|
||||
"model": "rerank-source",
|
||||
"query": "best document",
|
||||
"documents": ["alpha", {"text": "beta"}],
|
||||
"top_n": 1,
|
||||
"return_documents": true
|
||||
});
|
||||
let ctx = FormatContext::default().with_mapped_model("jina-reranker-v2-base-multilingual");
|
||||
|
||||
let converted = convert_request("openai:rerank", "jina:rerank", &body, &ctx)
|
||||
.expect("rerank request conversion should succeed");
|
||||
|
||||
assert_eq!(converted["model"], "jina-reranker-v2-base-multilingual");
|
||||
assert_eq!(converted["query"], "best document");
|
||||
assert_eq!(converted["documents"], json!(["alpha", {"text": "beta"}]));
|
||||
assert_eq!(converted["top_n"], 1);
|
||||
assert_eq!(converted["return_documents"], true);
|
||||
assert!(converted.get("messages").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rerank_registry_rejects_invalid_payloads() {
|
||||
let ctx = FormatContext::default();
|
||||
for body in [
|
||||
json!({"model": "rerank", "documents": ["alpha"]}),
|
||||
json!({"model": "rerank", "query": "q", "documents": []}),
|
||||
json!({"model": "rerank", "query": "q", "documents": [""]}),
|
||||
json!({"model": "rerank", "query": "q", "documents": ["alpha"], "top_n": 0}),
|
||||
] {
|
||||
assert!(convert_request("openai:rerank", "jina:rerank", &body, &ctx).is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registry_does_not_call_wire_specific_canonical_functions_directly() {
|
||||
let implementation = include_str!("registry.rs")
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
use crate::contracts::{
|
||||
CLAUDE_CHAT_STREAM_PLAN_KIND, CLAUDE_CHAT_SYNC_PLAN_KIND, CLAUDE_CLI_STREAM_PLAN_KIND,
|
||||
CLAUDE_CLI_SYNC_PLAN_KIND, GEMINI_CHAT_STREAM_PLAN_KIND, GEMINI_CHAT_SYNC_PLAN_KIND,
|
||||
GEMINI_CLI_STREAM_PLAN_KIND, GEMINI_CLI_SYNC_PLAN_KIND,
|
||||
GEMINI_CLI_STREAM_PLAN_KIND, GEMINI_CLI_SYNC_PLAN_KIND, OPENAI_EMBEDDING_SYNC_PLAN_KIND,
|
||||
OPENAI_RERANK_SYNC_PLAN_KIND,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
@@ -49,6 +50,20 @@ pub fn resolve_sync_spec(plan_kind: &str) -> Option<LocalSameFormatProviderSpec>
|
||||
family: LocalSameFormatProviderFamily::Gemini,
|
||||
require_streaming: false,
|
||||
}),
|
||||
OPENAI_EMBEDDING_SYNC_PLAN_KIND => Some(LocalSameFormatProviderSpec {
|
||||
api_format: "openai:embedding",
|
||||
decision_kind: OPENAI_EMBEDDING_SYNC_PLAN_KIND,
|
||||
report_kind: "openai_embedding_sync_success",
|
||||
family: LocalSameFormatProviderFamily::Standard,
|
||||
require_streaming: false,
|
||||
}),
|
||||
OPENAI_RERANK_SYNC_PLAN_KIND => Some(LocalSameFormatProviderSpec {
|
||||
api_format: "openai:rerank",
|
||||
decision_kind: OPENAI_RERANK_SYNC_PLAN_KIND,
|
||||
report_kind: "openai_rerank_sync_success",
|
||||
family: LocalSameFormatProviderFamily::Standard,
|
||||
require_streaming: false,
|
||||
}),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
@@ -106,4 +121,20 @@ mod tests {
|
||||
assert_eq!(spec.report_kind, "gemini_cli_stream_success");
|
||||
assert!(spec.require_streaming);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_openai_embedding_sync_same_format_spec() {
|
||||
let spec = resolve_sync_spec("openai_embedding_sync").expect("spec");
|
||||
assert_eq!(spec.api_format, "openai:embedding");
|
||||
assert_eq!(spec.report_kind, "openai_embedding_sync_success");
|
||||
assert!(!spec.require_streaming);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_openai_rerank_sync_same_format_spec() {
|
||||
let spec = resolve_sync_spec("openai_rerank_sync").expect("spec");
|
||||
assert_eq!(spec.api_format, "openai:rerank");
|
||||
assert_eq!(spec.report_kind, "openai_rerank_sync_success");
|
||||
assert!(!spec.require_streaming);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,12 +7,12 @@ use crate::contracts::{
|
||||
GEMINI_FILES_DOWNLOAD_PLAN_KIND, GEMINI_FILES_GET_PLAN_KIND, GEMINI_FILES_LIST_PLAN_KIND,
|
||||
GEMINI_FILES_UPLOAD_PLAN_KIND, GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND,
|
||||
GEMINI_VIDEO_CREATE_SYNC_PLAN_KIND, OPENAI_CHAT_STREAM_PLAN_KIND, OPENAI_CHAT_SYNC_PLAN_KIND,
|
||||
OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
|
||||
OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND, OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND,
|
||||
OPENAI_RESPONSES_STREAM_PLAN_KIND, OPENAI_RESPONSES_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND, OPENAI_VIDEO_CONTENT_PLAN_KIND,
|
||||
OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND, OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
|
||||
OPENAI_EMBEDDING_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
|
||||
OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND,
|
||||
OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND, OPENAI_RESPONSES_STREAM_PLAN_KIND,
|
||||
OPENAI_RESPONSES_SYNC_PLAN_KIND, OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_CONTENT_PLAN_KIND, OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND, OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
|
||||
};
|
||||
use crate::request::specialized::image::is_openai_image_stream_request;
|
||||
|
||||
@@ -173,6 +173,22 @@ pub fn resolve_execution_runtime_sync_plan_kind(
|
||||
return Some(OPENAI_CHAT_SYNC_PLAN_KIND);
|
||||
}
|
||||
|
||||
if route_family == Some("openai")
|
||||
&& route_kind == Some("embedding")
|
||||
&& *method == Method::POST
|
||||
&& path == "/v1/embeddings"
|
||||
{
|
||||
return Some(OPENAI_EMBEDDING_SYNC_PLAN_KIND);
|
||||
}
|
||||
|
||||
if route_family == Some("openai")
|
||||
&& route_kind == Some("rerank")
|
||||
&& *method == Method::POST
|
||||
&& path == "/v1/rerank"
|
||||
{
|
||||
return Some(OPENAI_RERANK_SYNC_PLAN_KIND);
|
||||
}
|
||||
|
||||
if route_family == Some("openai")
|
||||
&& route_kind == Some("image")
|
||||
&& *method == Method::POST
|
||||
@@ -327,6 +343,8 @@ pub fn supports_sync_execution_decision_kind(plan_kind: &str) -> bool {
|
||||
matches!(
|
||||
plan_kind,
|
||||
OPENAI_CHAT_SYNC_PLAN_KIND
|
||||
| OPENAI_EMBEDDING_SYNC_PLAN_KIND
|
||||
| OPENAI_RERANK_SYNC_PLAN_KIND
|
||||
| OPENAI_IMAGE_SYNC_PLAN_KIND
|
||||
| OPENAI_RESPONSES_SYNC_PLAN_KIND
|
||||
| OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND
|
||||
@@ -377,7 +395,8 @@ mod tests {
|
||||
CLAUDE_CHAT_STREAM_PLAN_KIND, CLAUDE_CHAT_SYNC_PLAN_KIND, CLAUDE_CLI_STREAM_PLAN_KIND,
|
||||
CLAUDE_CLI_SYNC_PLAN_KIND, GEMINI_CHAT_STREAM_PLAN_KIND, GEMINI_CHAT_SYNC_PLAN_KIND,
|
||||
GEMINI_CLI_STREAM_PLAN_KIND, GEMINI_CLI_SYNC_PLAN_KIND, OPENAI_CHAT_STREAM_PLAN_KIND,
|
||||
OPENAI_CHAT_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
|
||||
OPENAI_CHAT_SYNC_PLAN_KIND, OPENAI_EMBEDDING_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND,
|
||||
OPENAI_IMAGE_SYNC_PLAN_KIND, OPENAI_RERANK_SYNC_PLAN_KIND,
|
||||
OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND, OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND,
|
||||
OPENAI_RESPONSES_STREAM_PLAN_KIND, OPENAI_RESPONSES_SYNC_PLAN_KIND,
|
||||
};
|
||||
@@ -650,6 +669,40 @@ mod tests {
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_openai_embedding_sync_plan_kind() {
|
||||
assert_eq!(
|
||||
resolve_execution_runtime_sync_plan_kind(
|
||||
Some("ai_public"),
|
||||
Some("openai"),
|
||||
Some("embedding"),
|
||||
&Method::POST,
|
||||
"/v1/embeddings",
|
||||
),
|
||||
Some(OPENAI_EMBEDDING_SYNC_PLAN_KIND)
|
||||
);
|
||||
assert!(supports_sync_scheduler_decision_kind(
|
||||
OPENAI_EMBEDDING_SYNC_PLAN_KIND
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_openai_rerank_sync_plan_kind() {
|
||||
assert_eq!(
|
||||
resolve_execution_runtime_sync_plan_kind(
|
||||
Some("ai_public"),
|
||||
Some("openai"),
|
||||
Some("rerank"),
|
||||
&Method::POST,
|
||||
"/v1/rerank",
|
||||
),
|
||||
Some(OPENAI_RERANK_SYNC_PLAN_KIND)
|
||||
);
|
||||
assert!(supports_sync_scheduler_decision_kind(
|
||||
OPENAI_RERANK_SYNC_PLAN_KIND
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_openai_image_stream_plan_kind() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -8,6 +8,7 @@ pub enum LocalStandardSourceFamily {
|
||||
pub enum LocalStandardSourceMode {
|
||||
Chat,
|
||||
Cli,
|
||||
Embedding,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
|
||||
@@ -210,6 +210,12 @@ impl ProviderStreamParser {
|
||||
}
|
||||
FormatId::ClaudeMessages => Self::Claude(ClaudeProviderState::default()),
|
||||
FormatId::GeminiGenerateContent => Self::Gemini(GeminiProviderState::default()),
|
||||
FormatId::OpenAiEmbedding
|
||||
| FormatId::OpenAiRerank
|
||||
| FormatId::GeminiEmbedding
|
||||
| FormatId::JinaEmbedding
|
||||
| FormatId::JinaRerank
|
||||
| FormatId::DoubaoEmbedding => return None,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -294,6 +300,12 @@ impl ClientStreamEmitter {
|
||||
}
|
||||
FormatId::ClaudeMessages => Self::Claude(ClaudeClientEmitter::default()),
|
||||
FormatId::GeminiGenerateContent => Self::Gemini(GeminiClientEmitter::default()),
|
||||
FormatId::OpenAiEmbedding
|
||||
| FormatId::OpenAiRerank
|
||||
| FormatId::GeminiEmbedding
|
||||
| FormatId::JinaEmbedding
|
||||
| FormatId::JinaRerank
|
||||
| FormatId::DoubaoEmbedding => return None,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -352,6 +364,12 @@ fn parse_provider_error(
|
||||
}
|
||||
FormatId::ClaudeMessages => parse_claude_error(payload),
|
||||
FormatId::GeminiGenerateContent => parse_gemini_error(payload),
|
||||
FormatId::OpenAiEmbedding
|
||||
| FormatId::OpenAiRerank
|
||||
| FormatId::GeminiEmbedding
|
||||
| FormatId::JinaEmbedding
|
||||
| FormatId::JinaRerank
|
||||
| FormatId::DoubaoEmbedding => None,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,147 @@
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
const EMBEDDING_CAPABILITY: &str = "embedding";
|
||||
const EMBEDDING_API_FORMATS: &[&str] = &[
|
||||
"openai:embedding",
|
||||
"gemini:embedding",
|
||||
"jina:embedding",
|
||||
"doubao:embedding",
|
||||
"/v1/embeddings",
|
||||
"/jina/v1/embeddings",
|
||||
];
|
||||
|
||||
fn validate_optional_price(
|
||||
field_name: &str,
|
||||
value: Option<f64>,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
if value.is_some_and(|price| !price.is_finite() || price < 0.0) {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} must be a non-negative finite number"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_embedding_global_billing(
|
||||
default_price_per_request: Option<f64>,
|
||||
default_tiered_pricing: Option<&Value>,
|
||||
supported_capabilities: Option<&Value>,
|
||||
config: Option<&Value>,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
validate_optional_price(
|
||||
"global_models.default_price_per_request",
|
||||
default_price_per_request,
|
||||
)?;
|
||||
if !has_embedding_metadata(supported_capabilities, config) {
|
||||
return Ok(());
|
||||
}
|
||||
if has_request_or_input_token_pricing(default_price_per_request, default_tiered_pricing) {
|
||||
return Ok(());
|
||||
}
|
||||
Err(crate::DataLayerError::UnexpectedValue(
|
||||
"embedding global model requires default_price_per_request or default_tiered_pricing.tiers[].input_price_per_1m".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
fn validate_provider_model_pricing(
|
||||
price_per_request: Option<f64>,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
validate_optional_price("models.price_per_request", price_per_request)
|
||||
}
|
||||
|
||||
fn has_request_or_input_token_pricing(
|
||||
price_per_request: Option<f64>,
|
||||
tiered_pricing: Option<&Value>,
|
||||
) -> bool {
|
||||
price_per_request.is_some_and(|price| price.is_finite() && price >= 0.0)
|
||||
|| tiered_pricing
|
||||
.and_then(|value| value.get("tiers"))
|
||||
.and_then(Value::as_array)
|
||||
.is_some_and(|tiers| {
|
||||
tiers.iter().any(|tier| {
|
||||
tier.get("input_price_per_1m")
|
||||
.and_then(Value::as_f64)
|
||||
.is_some_and(|price| price.is_finite() && price >= 0.0)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn has_embedding_metadata(supported_capabilities: Option<&Value>, config: Option<&Value>) -> bool {
|
||||
supported_capabilities.is_some_and(value_contains_embedding_capability)
|
||||
|| config.is_some_and(value_contains_embedding_metadata)
|
||||
}
|
||||
|
||||
fn value_contains_embedding_capability(value: &Value) -> bool {
|
||||
match value {
|
||||
Value::String(value) => value.trim().eq_ignore_ascii_case(EMBEDDING_CAPABILITY),
|
||||
Value::Array(values) => values.iter().any(value_contains_embedding_capability),
|
||||
Value::Object(object) => {
|
||||
object
|
||||
.get(EMBEDDING_CAPABILITY)
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
|| [
|
||||
"capability",
|
||||
"model_type",
|
||||
"type",
|
||||
"task_type",
|
||||
"request_type",
|
||||
]
|
||||
.iter()
|
||||
.any(|key| {
|
||||
object
|
||||
.get(*key)
|
||||
.is_some_and(value_contains_embedding_capability)
|
||||
})
|
||||
|| ["capabilities", "supported_capabilities"]
|
||||
.iter()
|
||||
.any(|key| {
|
||||
object
|
||||
.get(*key)
|
||||
.is_some_and(value_contains_embedding_capability)
|
||||
})
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn value_contains_embedding_metadata(value: &Value) -> bool {
|
||||
match value {
|
||||
Value::String(value) => {
|
||||
value.trim().eq_ignore_ascii_case(EMBEDDING_CAPABILITY)
|
||||
|| is_known_embedding_api_format(value)
|
||||
}
|
||||
Value::Array(values) => values.iter().any(value_contains_embedding_metadata),
|
||||
Value::Object(object) => {
|
||||
value_contains_embedding_capability(value)
|
||||
|| ["api_format", "client_api_format", "provider_api_format"]
|
||||
.iter()
|
||||
.any(|key| {
|
||||
object
|
||||
.get(*key)
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(is_known_embedding_api_format)
|
||||
})
|
||||
|| ["api_formats", "client_api_formats", "provider_api_formats"]
|
||||
.iter()
|
||||
.any(|key| {
|
||||
object
|
||||
.get(*key)
|
||||
.is_some_and(value_contains_embedding_metadata)
|
||||
})
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn is_known_embedding_api_format(value: &str) -> bool {
|
||||
let normalized = value.trim().to_ascii_lowercase();
|
||||
EMBEDDING_API_FORMATS
|
||||
.iter()
|
||||
.any(|api_format| normalized == *api_format)
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredPublicGlobalModel {
|
||||
pub id: String,
|
||||
@@ -37,6 +178,12 @@ impl StoredPublicGlobalModel {
|
||||
"global_models.name is empty".to_string(),
|
||||
));
|
||||
}
|
||||
validate_embedding_global_billing(
|
||||
default_price_per_request,
|
||||
default_tiered_pricing.as_ref(),
|
||||
supported_capabilities.as_ref(),
|
||||
config.as_ref(),
|
||||
)?;
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
@@ -77,6 +224,7 @@ pub struct StoredPublicCatalogModel {
|
||||
pub supports_vision: Option<bool>,
|
||||
pub supports_function_calling: Option<bool>,
|
||||
pub supports_streaming: Option<bool>,
|
||||
pub supports_embedding: Option<bool>,
|
||||
pub is_active: bool,
|
||||
}
|
||||
|
||||
@@ -98,6 +246,7 @@ impl StoredPublicCatalogModel {
|
||||
supports_vision: Option<bool>,
|
||||
supports_function_calling: Option<bool>,
|
||||
supports_streaming: Option<bool>,
|
||||
supports_embedding: Option<bool>,
|
||||
is_active: bool,
|
||||
) -> Result<Self, crate::DataLayerError> {
|
||||
if id.trim().is_empty() {
|
||||
@@ -147,6 +296,7 @@ impl StoredPublicCatalogModel {
|
||||
supports_vision,
|
||||
supports_function_calling,
|
||||
supports_streaming,
|
||||
supports_embedding,
|
||||
is_active,
|
||||
})
|
||||
}
|
||||
@@ -223,6 +373,12 @@ impl StoredAdminGlobalModel {
|
||||
"global_models.display_name is empty".to_string(),
|
||||
));
|
||||
}
|
||||
validate_embedding_global_billing(
|
||||
default_price_per_request,
|
||||
default_tiered_pricing.as_ref(),
|
||||
supported_capabilities.as_ref(),
|
||||
config.as_ref(),
|
||||
)?;
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
@@ -273,6 +429,7 @@ pub struct StoredAdminProviderModel {
|
||||
pub global_model_display_name: Option<String>,
|
||||
pub global_model_default_price_per_request: Option<f64>,
|
||||
pub global_model_default_tiered_pricing: Option<Value>,
|
||||
pub global_model_supported_capabilities: Option<Value>,
|
||||
pub global_model_config: Option<Value>,
|
||||
}
|
||||
|
||||
@@ -300,6 +457,7 @@ impl StoredAdminProviderModel {
|
||||
global_model_display_name: Option<String>,
|
||||
global_model_default_price_per_request: Option<f64>,
|
||||
global_model_default_tiered_pricing: Option<Value>,
|
||||
global_model_supported_capabilities: Option<Value>,
|
||||
global_model_config: Option<Value>,
|
||||
) -> Result<Self, crate::DataLayerError> {
|
||||
if id.trim().is_empty() {
|
||||
@@ -322,6 +480,7 @@ impl StoredAdminProviderModel {
|
||||
"models.provider_model_name is empty".to_string(),
|
||||
));
|
||||
}
|
||||
validate_provider_model_pricing(price_per_request)?;
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
@@ -345,6 +504,7 @@ impl StoredAdminProviderModel {
|
||||
global_model_display_name,
|
||||
global_model_default_price_per_request,
|
||||
global_model_default_tiered_pricing,
|
||||
global_model_supported_capabilities,
|
||||
global_model_config,
|
||||
})
|
||||
}
|
||||
@@ -408,6 +568,7 @@ impl UpsertAdminProviderModelRecord {
|
||||
"models.provider_model_name is empty".to_string(),
|
||||
));
|
||||
}
|
||||
validate_provider_model_pricing(price_per_request)?;
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
@@ -468,6 +629,12 @@ impl CreateAdminGlobalModelRecord {
|
||||
"global_models.display_name is empty".to_string(),
|
||||
));
|
||||
}
|
||||
validate_embedding_global_billing(
|
||||
default_price_per_request,
|
||||
default_tiered_pricing.as_ref(),
|
||||
supported_capabilities.as_ref(),
|
||||
config.as_ref(),
|
||||
)?;
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
@@ -514,6 +681,12 @@ impl UpdateAdminGlobalModelRecord {
|
||||
"global_models.display_name is empty".to_string(),
|
||||
));
|
||||
}
|
||||
validate_embedding_global_billing(
|
||||
default_price_per_request,
|
||||
default_tiered_pricing.as_ref(),
|
||||
supported_capabilities.as_ref(),
|
||||
config.as_ref(),
|
||||
)?;
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
@@ -695,3 +868,99 @@ pub trait GlobalModelWriteRepository: Send + Sync {
|
||||
global_model_id: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{CreateAdminGlobalModelRecord, UpsertAdminProviderModelRecord};
|
||||
|
||||
#[test]
|
||||
fn embedding_missing_billing_config_rejected() {
|
||||
let err = CreateAdminGlobalModelRecord::new(
|
||||
"gm-embedding".to_string(),
|
||||
"text-embedding-3-small".to_string(),
|
||||
"Text Embedding 3 Small".to_string(),
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
Some(json!(["embedding"])),
|
||||
None,
|
||||
)
|
||||
.expect_err("embedding model without explicit billing should be rejected");
|
||||
|
||||
assert!(err.to_string().contains(
|
||||
"embedding global model requires default_price_per_request or default_tiered_pricing.tiers[].input_price_per_1m"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_api_format_config_requires_billing_config() {
|
||||
let err = CreateAdminGlobalModelRecord::new(
|
||||
"gm-embedding".to_string(),
|
||||
"jina-embeddings-v3".to_string(),
|
||||
"Jina Embeddings v3".to_string(),
|
||||
true,
|
||||
None,
|
||||
Some(json!({"tiers": []})),
|
||||
None,
|
||||
Some(json!({"api_formats": ["jina:embedding"]})),
|
||||
)
|
||||
.expect_err("embedding API format without price should be rejected");
|
||||
|
||||
assert!(err.to_string().contains("embedding global model requires"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_input_token_or_request_pricing_is_accepted() {
|
||||
CreateAdminGlobalModelRecord::new(
|
||||
"gm-embedding-input".to_string(),
|
||||
"text-embedding-3-small".to_string(),
|
||||
"Text Embedding 3 Small".to_string(),
|
||||
true,
|
||||
None,
|
||||
Some(json!({"tiers":[{"up_to":null,"input_price_per_1m":0.02}]})),
|
||||
Some(json!(["embedding"])),
|
||||
Some(json!({"dimensions": 1536})),
|
||||
)
|
||||
.expect("input-token pricing should satisfy embedding billing");
|
||||
|
||||
CreateAdminGlobalModelRecord::new(
|
||||
"gm-embedding-request".to_string(),
|
||||
"custom-embedding".to_string(),
|
||||
"Custom Embedding".to_string(),
|
||||
true,
|
||||
Some(0.0),
|
||||
None,
|
||||
Some(json!(["embedding"])),
|
||||
None,
|
||||
)
|
||||
.expect("explicit request pricing should satisfy embedding billing");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_model_negative_request_price_rejected() {
|
||||
let err = UpsertAdminProviderModelRecord::new(
|
||||
"model-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"global-model-1".to_string(),
|
||||
"text-embedding-3-small".to_string(),
|
||||
None,
|
||||
Some(-0.01),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
true,
|
||||
None,
|
||||
)
|
||||
.expect_err("negative provider model request price should be rejected");
|
||||
|
||||
assert!(err
|
||||
.to_string()
|
||||
.contains("models.price_per_request must be a non-negative finite number"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -497,6 +497,7 @@ impl GlobalModelWriteRepository for InMemoryGlobalModelReadRepository {
|
||||
Some(global_model.display_name.clone()),
|
||||
global_model.default_price_per_request,
|
||||
global_model.default_tiered_pricing.clone(),
|
||||
global_model.supported_capabilities.clone(),
|
||||
global_model.config.clone(),
|
||||
)?;
|
||||
self.admin_provider_model_items
|
||||
@@ -542,6 +543,7 @@ impl GlobalModelWriteRepository for InMemoryGlobalModelReadRepository {
|
||||
existing.global_model_display_name = Some(global_model.display_name.clone());
|
||||
existing.global_model_default_price_per_request = global_model.default_price_per_request;
|
||||
existing.global_model_default_tiered_pricing = global_model.default_tiered_pricing.clone();
|
||||
existing.global_model_supported_capabilities = global_model.supported_capabilities.clone();
|
||||
existing.global_model_config = global_model.config.clone();
|
||||
Ok(Some(existing.clone()))
|
||||
}
|
||||
@@ -639,8 +641,9 @@ mod tests {
|
||||
|
||||
use super::InMemoryGlobalModelReadRepository;
|
||||
use crate::repository::global_models::{
|
||||
GlobalModelReadRepository, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
|
||||
PublicGlobalModelQuery, StoredPublicCatalogModel, StoredPublicGlobalModel,
|
||||
CreateAdminGlobalModelRecord, GlobalModelReadRepository, GlobalModelWriteRepository,
|
||||
PublicCatalogModelListQuery, PublicCatalogModelSearchQuery, PublicGlobalModelQuery,
|
||||
StoredPublicCatalogModel, StoredPublicGlobalModel,
|
||||
};
|
||||
|
||||
fn sample_model(
|
||||
@@ -687,11 +690,83 @@ mod tests {
|
||||
Some(true),
|
||||
Some(true),
|
||||
Some(true),
|
||||
Some(false),
|
||||
true,
|
||||
)
|
||||
.expect("public catalog model should build")
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn embedding_model_metadata_roundtrip() {
|
||||
let repository =
|
||||
InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new());
|
||||
let record = CreateAdminGlobalModelRecord::new(
|
||||
"gm-embedding".to_string(),
|
||||
"text-embedding-3-small".to_string(),
|
||||
"Text Embedding 3 Small".to_string(),
|
||||
true,
|
||||
None,
|
||||
Some(json!({"tiers":[{"up_to":null,"input_price_per_1m":0.02}]})),
|
||||
Some(json!(["embedding"])),
|
||||
Some(json!({
|
||||
"api_formats": ["openai:embedding"],
|
||||
"dimensions": 1536
|
||||
})),
|
||||
)
|
||||
.expect("embedding global model should validate");
|
||||
|
||||
repository
|
||||
.create_admin_global_model(&record)
|
||||
.await
|
||||
.expect("embedding global model should persist")
|
||||
.expect("embedding global model should be returned");
|
||||
|
||||
let stored = repository
|
||||
.get_admin_global_model_by_name("text-embedding-3-small")
|
||||
.await
|
||||
.expect("embedding global model should read")
|
||||
.expect("embedding global model should exist");
|
||||
|
||||
assert_eq!(stored.supported_capabilities, Some(json!(["embedding"])));
|
||||
assert_eq!(
|
||||
stored
|
||||
.config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("dimensions")),
|
||||
Some(&json!(1536))
|
||||
);
|
||||
assert_eq!(
|
||||
stored
|
||||
.default_tiered_pricing
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("tiers"))
|
||||
.and_then(serde_json::Value::as_array)
|
||||
.and_then(|tiers| tiers.first())
|
||||
.and_then(|tier| tier.get("input_price_per_1m"))
|
||||
.and_then(serde_json::Value::as_f64),
|
||||
Some(0.02)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn embedding_missing_billing_config_rejected() {
|
||||
let error = CreateAdminGlobalModelRecord::new(
|
||||
"gm-embedding".to_string(),
|
||||
"text-embedding-3-small".to_string(),
|
||||
"Text Embedding 3 Small".to_string(),
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
Some(json!(["embedding"])),
|
||||
None,
|
||||
)
|
||||
.expect_err("embedding metadata without billing should fail closed");
|
||||
|
||||
assert!(error
|
||||
.to_string()
|
||||
.contains("embedding global model requires"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn defaults_to_active_models_only() {
|
||||
let repository = InMemoryGlobalModelReadRepository::seed(vec![
|
||||
@@ -791,6 +866,52 @@ mod tests {
|
||||
assert_eq!(items[0].name, "gpt-5");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn public_catalog_preserves_embedding_capability_without_contaminating_chat_models() {
|
||||
let mut embedding_model = sample_public_catalog_model(
|
||||
"model-embedding",
|
||||
"provider-openai",
|
||||
"openai",
|
||||
"text-embedding-3-small",
|
||||
"text-embedding-3-small",
|
||||
"Text Embedding 3 Small",
|
||||
);
|
||||
embedding_model.supports_embedding = Some(true);
|
||||
embedding_model.supports_streaming = Some(false);
|
||||
let chat_model = sample_public_catalog_model(
|
||||
"model-chat",
|
||||
"provider-openai",
|
||||
"openai",
|
||||
"gpt-5-upstream",
|
||||
"gpt-5",
|
||||
"GPT 5",
|
||||
);
|
||||
let repository =
|
||||
InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new())
|
||||
.with_public_catalog_models(vec![embedding_model, chat_model]);
|
||||
|
||||
let items = repository
|
||||
.list_public_catalog_models(&PublicCatalogModelListQuery {
|
||||
provider_id: Some("provider-openai".to_string()),
|
||||
offset: 0,
|
||||
limit: 50,
|
||||
})
|
||||
.await
|
||||
.expect("catalog should list");
|
||||
|
||||
let embedding = items
|
||||
.iter()
|
||||
.find(|item| item.name == "text-embedding-3-small")
|
||||
.expect("embedding model should be listed");
|
||||
let chat = items
|
||||
.iter()
|
||||
.find(|item| item.name == "gpt-5")
|
||||
.expect("chat model should be listed");
|
||||
assert_eq!(embedding.supports_embedding, Some(true));
|
||||
assert_eq!(embedding.supports_streaming, Some(false));
|
||||
assert_eq!(chat.supports_embedding, Some(false));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn searches_public_catalog_models_by_provider_and_display_name() {
|
||||
let repository =
|
||||
|
||||
@@ -60,6 +60,27 @@ SELECT
|
||||
COALESCE(m.supports_vision, CAST(gm.config->>'vision' AS BOOLEAN), FALSE) AS supports_vision,
|
||||
COALESCE(m.supports_function_calling, CAST(gm.config->>'function_calling' AS BOOLEAN), FALSE) AS supports_function_calling,
|
||||
COALESCE(m.supports_streaming, CAST(gm.config->>'streaming' AS BOOLEAN), TRUE) AS supports_streaming,
|
||||
(
|
||||
COALESCE(gm.supported_capabilities::jsonb @> '["embedding"]'::jsonb, FALSE)
|
||||
OR LOWER(COALESCE(gm.config->>'embedding', 'false')) = 'true'
|
||||
OR LOWER(COALESCE(gm.config->>'model_type', '')) = 'embedding'
|
||||
OR LOWER(COALESCE(gm.config->>'type', '')) = 'embedding'
|
||||
OR COALESCE(gm.config->'capabilities' @> '["embedding"]'::jsonb, FALSE)
|
||||
OR COALESCE(gm.config->'supported_capabilities' @> '["embedding"]'::jsonb, FALSE)
|
||||
OR COALESCE(gm.config->'api_formats' @> '["openai:embedding"]'::jsonb, FALSE)
|
||||
OR COALESCE(gm.config->'api_formats' @> '["jina:embedding"]'::jsonb, FALSE)
|
||||
OR COALESCE(gm.config->'api_formats' @> '["gemini:embedding"]'::jsonb, FALSE)
|
||||
OR COALESCE(gm.config->'api_formats' @> '["doubao:embedding"]'::jsonb, FALSE)
|
||||
OR LOWER(COALESCE(m.config->>'embedding', 'false')) = 'true'
|
||||
OR LOWER(COALESCE(m.config->>'model_type', '')) = 'embedding'
|
||||
OR LOWER(COALESCE(m.config->>'type', '')) = 'embedding'
|
||||
OR COALESCE(m.config::jsonb->'capabilities' @> '["embedding"]'::jsonb, FALSE)
|
||||
OR COALESCE(m.config::jsonb->'supported_capabilities' @> '["embedding"]'::jsonb, FALSE)
|
||||
OR COALESCE(m.config::jsonb->'api_formats' @> '["openai:embedding"]'::jsonb, FALSE)
|
||||
OR COALESCE(m.config::jsonb->'api_formats' @> '["jina:embedding"]'::jsonb, FALSE)
|
||||
OR COALESCE(m.config::jsonb->'api_formats' @> '["gemini:embedding"]'::jsonb, FALSE)
|
||||
OR COALESCE(m.config::jsonb->'api_formats' @> '["doubao:embedding"]'::jsonb, FALSE)
|
||||
) AS supports_embedding,
|
||||
m.is_active
|
||||
FROM models m
|
||||
JOIN providers p ON p.id = m.provider_id
|
||||
@@ -98,6 +119,7 @@ SELECT
|
||||
gm.display_name AS global_model_display_name,
|
||||
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS global_model_default_price_per_request,
|
||||
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
|
||||
gm.supported_capabilities AS global_model_supported_capabilities,
|
||||
gm.config AS global_model_config
|
||||
FROM models m
|
||||
LEFT JOIN global_models gm ON gm.id = m.global_model_id
|
||||
@@ -302,6 +324,7 @@ SELECT
|
||||
gm.display_name AS global_model_display_name,
|
||||
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS global_model_default_price_per_request,
|
||||
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
|
||||
gm.supported_capabilities AS global_model_supported_capabilities,
|
||||
gm.config AS global_model_config
|
||||
FROM models m
|
||||
LEFT JOIN global_models gm ON gm.id = m.global_model_id
|
||||
@@ -348,6 +371,7 @@ SELECT
|
||||
gm.display_name AS global_model_display_name,
|
||||
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS global_model_default_price_per_request,
|
||||
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
|
||||
gm.supported_capabilities AS global_model_supported_capabilities,
|
||||
gm.config AS global_model_config
|
||||
FROM models m
|
||||
JOIN global_models gm ON gm.id = m.global_model_id
|
||||
@@ -487,6 +511,7 @@ SELECT
|
||||
gm.display_name AS global_model_display_name,
|
||||
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS global_model_default_price_per_request,
|
||||
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
|
||||
gm.supported_capabilities AS global_model_supported_capabilities,
|
||||
gm.config AS global_model_config
|
||||
FROM models m
|
||||
LEFT JOIN global_models gm ON gm.id = m.global_model_id
|
||||
@@ -1020,7 +1045,7 @@ fn apply_public_catalog_model_filters(
|
||||
provider_id: Option<&str>,
|
||||
search: Option<&str>,
|
||||
) {
|
||||
builder.push(" WHERE m.is_active = TRUE AND p.is_active = TRUE");
|
||||
builder.push(" WHERE m.is_active = TRUE AND COALESCE(m.is_available, TRUE) = TRUE AND p.is_active = TRUE AND COALESCE(gm.is_active, TRUE) = TRUE");
|
||||
|
||||
if let Some(provider_id) = provider_id.map(str::trim).filter(|value| !value.is_empty()) {
|
||||
builder
|
||||
@@ -1060,6 +1085,7 @@ fn map_public_catalog_model_row(row: &PgRow) -> Result<StoredPublicCatalogModel,
|
||||
row.try_get("supports_function_calling")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("supports_streaming").map_postgres_err()?,
|
||||
row.try_get("supports_embedding").map_postgres_err()?,
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
@@ -1101,6 +1127,8 @@ fn map_admin_provider_model_row(row: &PgRow) -> Result<StoredAdminProviderModel,
|
||||
.map_postgres_err()?,
|
||||
row.try_get("global_model_default_tiered_pricing")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("global_model_supported_capabilities")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("global_model_config").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
@@ -1177,9 +1205,42 @@ fn map_provider_active_global_model_row(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::SqlxGlobalModelReadRepository;
|
||||
use super::{SqlxGlobalModelReadRepository, LIST_ADMIN_PROVIDER_MODELS_PREFIX};
|
||||
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
|
||||
const ADMIN_PROVIDER_MODEL_REQUIRED_COLUMNS: &[&str] = &[
|
||||
"global_model_default_tiered_pricing",
|
||||
"global_model_supported_capabilities",
|
||||
"global_model_config",
|
||||
];
|
||||
|
||||
fn assert_admin_provider_model_projection_has_required_columns(sql: &str) {
|
||||
for column in ADMIN_PROVIDER_MODEL_REQUIRED_COLUMNS {
|
||||
assert!(
|
||||
sql.contains(column),
|
||||
"admin provider model SQL projection should include {column}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn admin_provider_model_sql_projections_include_supported_capabilities() {
|
||||
assert_admin_provider_model_projection_has_required_columns(
|
||||
LIST_ADMIN_PROVIDER_MODELS_PREFIX,
|
||||
);
|
||||
assert_admin_provider_model_projection_has_required_columns(include_str!("sql.rs"));
|
||||
let supported_capabilities_projection = format!(
|
||||
"{} AS {}",
|
||||
"gm.supported_capabilities", "global_model_supported_capabilities"
|
||||
);
|
||||
assert_eq!(
|
||||
include_str!("sql.rs")
|
||||
.matches(&supported_capabilities_projection)
|
||||
.count(),
|
||||
4
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_constructs_from_lazy_pool() {
|
||||
let factory = PostgresPoolFactory::new(PostgresPoolConfig {
|
||||
|
||||
@@ -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