use std::sync::Arc;
use anyhow::Result;
use serde::{Deserialize, Serialize};
use crate::shared::api::{EmbedRole, Embedder};
const INSTRUCT_QUERY_PREFIX: &str =
"Instruct: Given a query, retrieve passages that answer it\nQuery: ";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum EmbedConvention {
#[default]
None,
E5,
E5Instruct,
}
impl EmbedConvention {
pub const ALL: [Self; 3] = [Self::None, Self::E5, Self::E5Instruct];
pub fn id(self) -> &'static str {
match self {
Self::None => "none",
Self::E5 => "e5",
Self::E5Instruct => "e5-instruct",
}
}
pub fn prefix(self, role: EmbedRole) -> &'static str {
match (self, role) {
(Self::None, _) => "",
(Self::E5, EmbedRole::Query) => "query: ",
(Self::E5, EmbedRole::Passage) => "passage: ",
(Self::E5Instruct, EmbedRole::Query) => INSTRUCT_QUERY_PREFIX,
(Self::E5Instruct, EmbedRole::Passage) => "",
}
}
pub fn cycle(self, dir: i32) -> Self {
let i = Self::ALL.iter().position(|c| *c == self).unwrap_or(0) as i32;
let n = Self::ALL.len() as i32;
Self::ALL[(i + dir).rem_euclid(n) as usize]
}
pub fn suggested_for(model_id: &str) -> Option<Self> {
let name = model_id.to_ascii_lowercase();
if !name.contains("e5") {
return None;
}
Some(if name.contains("instruct") {
Self::E5Instruct
} else {
Self::E5
})
}
}
pub struct PrefixedEmbedder {
inner: Arc<dyn Embedder>,
convention: EmbedConvention,
}
impl PrefixedEmbedder {
pub fn new(inner: Arc<dyn Embedder>, convention: EmbedConvention) -> Self {
Self { inner, convention }
}
}
#[async_trait::async_trait]
impl Embedder for PrefixedEmbedder {
async fn embed(&self, texts: Vec<String>, role: EmbedRole) -> Result<Vec<Vec<f32>>> {
let prefix = self.convention.prefix(role);
if prefix.is_empty() {
return self.inner.embed(texts, role).await;
}
let marked = texts.into_iter().map(|t| format!("{prefix}{t}")).collect();
self.inner.embed(marked, role).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex;
#[derive(Default)]
struct Recorder {
seen: Mutex<Vec<(String, EmbedRole)>>,
}
#[async_trait::async_trait]
impl Embedder for Recorder {
async fn embed(&self, texts: Vec<String>, role: EmbedRole) -> Result<Vec<Vec<f32>>> {
let mut seen = self.seen.lock().unwrap();
for t in &texts {
seen.push((t.clone(), role));
}
Ok(texts.iter().map(|_| vec![1.0, 0.0]).collect())
}
}
async fn sent(convention: EmbedConvention, role: EmbedRole) -> Vec<String> {
let rec = Arc::new(Recorder::default());
let e = PrefixedEmbedder::new(rec.clone(), convention);
e.embed(vec!["сколько стоит?".into(), "how much?".into()], role)
.await
.unwrap();
let seen = rec.seen.lock().unwrap();
seen.iter().map(|(t, _)| t.clone()).collect()
}
#[tokio::test]
async fn none_passes_text_through_byte_for_byte() {
for role in [EmbedRole::Query, EmbedRole::Passage] {
assert_eq!(
sent(EmbedConvention::None, role).await,
["сколько стоит?", "how much?"],
"{role:?}"
);
}
}
#[tokio::test]
async fn e5_marks_both_sides_differently() {
assert_eq!(
sent(EmbedConvention::E5, EmbedRole::Query).await[0],
"query: сколько стоит?"
);
assert_eq!(
sent(EmbedConvention::E5, EmbedRole::Passage).await[0],
"passage: сколько стоит?"
);
}
#[tokio::test]
async fn e5_instruct_marks_only_the_query() {
let q = sent(EmbedConvention::E5Instruct, EmbedRole::Query).await;
assert!(q[0].starts_with("Instruct: "), "{q:?}");
assert!(q[0].ends_with("\nQuery: сколько стоит?"), "{q:?}");
assert_eq!(
sent(EmbedConvention::E5Instruct, EmbedRole::Passage).await,
["сколько стоит?", "how much?"],
"the instruct variant's passages are bare"
);
}
#[tokio::test]
async fn the_role_reaches_the_inner_embedder_unchanged() {
let rec = Arc::new(Recorder::default());
let e = PrefixedEmbedder::new(rec.clone(), EmbedConvention::E5);
e.embed(vec!["x".into()], EmbedRole::Query).await.unwrap();
e.embed(vec!["y".into()], EmbedRole::Passage).await.unwrap();
let seen = rec.seen.lock().unwrap();
assert_eq!(seen[0].1, EmbedRole::Query);
assert_eq!(seen[1].1, EmbedRole::Passage);
}
#[test]
fn every_convention_has_a_stable_id_and_all_are_listed() {
assert_eq!(EmbedConvention::None.id(), "none");
assert_eq!(EmbedConvention::E5.id(), "e5");
assert_eq!(EmbedConvention::E5Instruct.id(), "e5-instruct");
assert_eq!(EmbedConvention::ALL.len(), 3);
assert_eq!(EmbedConvention::default(), EmbedConvention::None);
}
#[test]
fn only_none_is_a_no_op() {
for c in EmbedConvention::ALL {
let touches = [EmbedRole::Query, EmbedRole::Passage]
.iter()
.any(|r| !c.prefix(*r).is_empty());
assert_eq!(touches, c != EmbedConvention::None, "{c:?}");
}
}
#[test]
fn cycle_wraps_in_both_directions() {
assert_eq!(EmbedConvention::None.cycle(1), EmbedConvention::E5);
assert_eq!(EmbedConvention::E5Instruct.cycle(1), EmbedConvention::None);
assert_eq!(EmbedConvention::None.cycle(-1), EmbedConvention::E5Instruct);
}
#[test]
fn suggestion_recognises_the_e5_family_and_nothing_else() {
assert_eq!(
EmbedConvention::suggested_for(r"D:\LLM\GGUF\multilingual-e5-large-instruct-q8_0.gguf"),
Some(EmbedConvention::E5Instruct)
);
assert_eq!(
EmbedConvention::suggested_for("intfloat/multilingual-e5-large"),
Some(EmbedConvention::E5)
);
assert_eq!(EmbedConvention::suggested_for("bge-m3-Q8_0.gguf"), None);
assert_eq!(EmbedConvention::suggested_for(""), None);
}
#[test]
fn the_instruct_query_marker_has_the_documented_shape() {
let p = EmbedConvention::E5Instruct.prefix(EmbedRole::Query);
assert_eq!(p, INSTRUCT_QUERY_PREFIX);
assert!(
p.starts_with("Instruct: ") && p.ends_with("\nQuery: "),
"{p}"
);
}
#[tokio::test]
#[ignore = "requires two live embedding servers (MINDFORK_EMBED_URL, MINDFORK_EMBED_URL_ALT)"]
async fn conventions_behave_as_measured_live() {
const BGE: (&str, &str) = ("MINDFORK_EMBED_URL", "MINDFORK_EMBED_KEY");
const E5: (&str, &str) = ("MINDFORK_EMBED_URL_ALT", "MINDFORK_EMBED_KEY_ALT");
if crate::shared::api::live_client(BGE.0, BGE.1).is_none()
|| crate::shared::api::live_client(E5.0, E5.1).is_none()
{
eprintln!("skip: MINDFORK_EMBED_URL / MINDFORK_EMBED_URL_ALT not set");
return;
}
const QUERY: &str = "какой внутренний код сборки проекта?";
const RELEVANT: &str =
"Внутренний код сборки проекта — ZARYA-7719, он указывается в отчётах о релизе.";
const IRRELEVANT: &str = "Чугунной сковороде нужна прокалка перед первым использованием.";
async fn margin(server: (&str, &str), c: EmbedConvention) -> f32 {
let raw: Arc<dyn Embedder> =
Arc::new(crate::shared::api::live_client(server.0, server.1).unwrap());
let e = PrefixedEmbedder::new(raw, c);
let q = e
.embed(vec![QUERY.into()], EmbedRole::Query)
.await
.unwrap()
.remove(0);
let docs = e
.embed(vec![RELEVANT.into(), IRRELEVANT.into()], EmbedRole::Passage)
.await
.unwrap();
cos(&q, &docs[0]) - cos(&q, &docs[1])
}
fn cos(a: &[f32], b: &[f32]) -> f32 {
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
dot / (na * nb)
}
let e5_none = margin(E5, EmbedConvention::None).await;
let e5_own = margin(E5, EmbedConvention::E5Instruct).await;
let bge_none = margin(BGE, EmbedConvention::None).await;
let bge_wrong = margin(BGE, EmbedConvention::E5Instruct).await;
eprintln!(
"e5: none={e5_none:.4} own={e5_own:.4} | bge: none={bge_none:.4} wrong={bge_wrong:.4}"
);
assert!(
e5_own > e5_none,
"e5's own convention must separate better: {e5_none:.4} -> {e5_own:.4}"
);
assert!(
bge_none > bge_wrong,
"and the wrong convention must hurt bge-m3, which is why `none` is the \
default and nothing is auto-applied: {bge_none:.4} -> {bge_wrong:.4}"
);
}
}