mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
feat: add embedding and rerank support
This commit is contained in:
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user