Merge remote-tracking branch 'origin/pr/593'

This commit is contained in:
elky
2026-06-02 21:40:59 +08:00
56 changed files with 1605 additions and 42 deletions
@@ -246,6 +246,19 @@ pub enum CanonicalEmbeddingInput {
StringArray(Vec<String>),
TokenArray(Vec<i64>),
TokenArrayArray(Vec<Vec<i64>>),
Multimodal(Vec<CanonicalEmbeddingContent>),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CanonicalEmbeddingContent {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub text: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub image: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub video: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub multi_images: Option<Vec<String>>,
}
impl CanonicalEmbeddingInput {
@@ -257,6 +270,9 @@ impl CanonicalEmbeddingInput {
}
Self::TokenArray(values) => values.is_empty(),
Self::TokenArrayArray(values) => values.is_empty() || values.iter().any(Vec::is_empty),
Self::Multimodal(values) => {
values.is_empty() || values.iter().any(CanonicalEmbeddingContent::is_empty)
}
}
}
@@ -264,11 +280,47 @@ impl CanonicalEmbeddingInput {
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,
Self::TokenArray(_) | Self::TokenArrayArray(_) | Self::Multimodal(_) => None,
}
}
}
impl CanonicalEmbeddingContent {
pub(crate) fn is_empty(&self) -> bool {
let text_empty = self
.text
.as_ref()
.is_some_and(|value| value.trim().is_empty());
let image_empty = self
.image
.as_ref()
.is_some_and(|value| value.trim().is_empty());
let video_empty = self
.video
.as_ref()
.is_some_and(|value| value.trim().is_empty());
let multi_images_empty = self.multi_images.as_ref().is_some_and(|values| {
values.is_empty() || values.iter().any(|value| value.trim().is_empty())
});
let has_any = self
.text
.as_ref()
.is_some_and(|value| !value.trim().is_empty())
|| self
.image
.as_ref()
.is_some_and(|value| !value.trim().is_empty())
|| self
.video
.as_ref()
.is_some_and(|value| !value.trim().is_empty())
|| self.multi_images.as_ref().is_some_and(|values| {
!values.is_empty() && values.iter().all(|value| !value.trim().is_empty())
});
!has_any || text_empty || image_empty || video_empty || multi_images_empty
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CanonicalEmbeddingRequest {
pub input: CanonicalEmbeddingInput,
@@ -280,6 +332,8 @@ pub struct CanonicalEmbeddingRequest {
pub task: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub parameters: Option<Map<String, Value>>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub extensions: BTreeMap<String, Value>,
}
@@ -502,6 +556,7 @@ pub(crate) fn canonical_to_embedding_request(
"jina" => crate::formats::jina::embedding::request::to(canonical, &ctx),
"gemini" => crate::formats::gemini::embedding::request::to(canonical, &ctx),
"doubao" => crate::formats::doubao::embedding::request::to(canonical, &ctx),
"aliyun" => crate::formats::aliyun::embedding::request::to(canonical, &ctx),
_ => None,
}
}
@@ -699,6 +754,7 @@ pub fn from_embedding_to_canonical_response(
}
"jina" => crate::formats::openai::embedding::response::from_namespace(body_json, "jina"),
"gemini" => crate::formats::gemini::embedding::response::from(body_json),
"aliyun" => crate::formats::aliyun::embedding::response::from(body_json),
_ => None,
}
}
@@ -5244,8 +5300,8 @@ 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, CanonicalEmbedding, CanonicalEmbeddingInput,
CanonicalEmbeddingRequest, CanonicalRole, CanonicalUsage,
CanonicalContentBlock, CanonicalEmbedding, CanonicalEmbeddingContent,
CanonicalEmbeddingInput, CanonicalEmbeddingRequest, CanonicalRole, CanonicalUsage,
};
use serde_json::{json, Value};
@@ -5301,6 +5357,44 @@ mod tests {
"nested token array",
CanonicalEmbeddingInput::TokenArrayArray(vec![vec![1, 2], vec![3, 4]]),
),
(
json!([
{"text": "white running shoes"},
{"image": "https://example.com/shoe.png"},
{"video": "https://example.com/demo.mp4"},
{"multi_images": ["https://example.com/a.png", "https://example.com/b.png"]}
]),
"multimodal array",
CanonicalEmbeddingInput::Multimodal(vec![
CanonicalEmbeddingContent {
text: Some("white running shoes".to_string()),
image: None,
video: None,
multi_images: None,
},
CanonicalEmbeddingContent {
text: None,
image: Some("https://example.com/shoe.png".to_string()),
video: None,
multi_images: None,
},
CanonicalEmbeddingContent {
text: None,
image: None,
video: Some("https://example.com/demo.mp4".to_string()),
multi_images: None,
},
CanonicalEmbeddingContent {
text: None,
image: None,
video: None,
multi_images: Some(vec![
"https://example.com/a.png".to_string(),
"https://example.com/b.png".to_string(),
]),
},
]),
),
];
for (input, label, expected_input) in cases {
@@ -5327,6 +5421,9 @@ mod tests {
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": "text-embedding-3-small", "input": [{"image": " "}]}),
json!({"model": "text-embedding-3-small", "input": [{"multi_images": []}]}),
json!({"model": "text-embedding-3-small", "input": ["hello", {"image": "https://example.com/a.png"}]}),
json!({"model": "", "input": "hello"}),
json!({"input": "hello"}),
json!({"model": "text-embedding-3-small", "messages": []}),
@@ -5396,6 +5493,7 @@ mod tests {
dimensions: Some(2),
task: None,
user: None,
parameters: None,
extensions: Default::default(),
}),
..Default::default()
@@ -5441,6 +5539,7 @@ mod tests {
dimensions: Some(1536),
task: Some("retrieval.passage".to_string()),
user: Some("user-1".to_string()),
parameters: None,
extensions: Default::default(),
}),
..Default::default()
@@ -5493,6 +5592,7 @@ mod tests {
dimensions: None,
task: None,
user: None,
parameters: None,
extensions: Default::default(),
}),
..Default::default()