use crate::embeddings::embed::{EmbedError, TextEmbedder};
use crate::embeddings::{Embed, Embedding};
use crate::error::ProviderError;
use super::EmbeddingsBuilder;
use crate::driver::{Exchange, Local, Model, Opened, Opening, Step, Transport};
use crate::test_utils::MockEmbeddings;
use crate::wire::Capabilities;
fn batches(max_documents: usize) -> Local<crate::operation::Embedding> {
Local::new("mock").with_capabilities(Capabilities::embedding(max_documents, 10))
}
#[derive(Debug)]
struct NTexts {
doc: usize,
n: usize,
}
impl NTexts {
fn new(doc: usize, n: usize) -> Self {
Self { doc, n }
}
fn expected(&self) -> Vec<String> {
(0..self.n).map(|i| format!("d{}t{i}", self.doc)).collect()
}
}
impl Embed for NTexts {
fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
for i in 0..self.n {
embedder.embed(format!("d{}t{i}", self.doc));
}
Ok(())
}
}
fn returned(embeddings: &[Embedding]) -> Vec<String> {
embeddings
.iter()
.map(|embedding| embedding.document.clone())
.collect()
}
#[tokio::test]
async fn test_build_rejects_a_document_that_embeds_no_text() {
let error = EmbeddingsBuilder::new(MockEmbeddings::model())
.document(NTexts::new(0, 0))
.unwrap()
.build()
.await
.expect_err("a document with no texts has no embeddings");
assert!(
matches!(error, ProviderError::Response(_)),
"unexpected error variant: {error:?}"
);
assert!(
error.to_string().contains("document 0 produced no text"),
"error should name the offending document: {error}"
);
}
#[tokio::test]
async fn test_build_names_the_document_that_embeds_no_text() {
let error = EmbeddingsBuilder::new(MockEmbeddings::model())
.documents(vec![
NTexts::new(0, 2),
NTexts::new(1, 2),
NTexts::new(2, 0),
])
.unwrap()
.build()
.await
.expect_err("a document with no texts has no embeddings");
assert!(
error.to_string().contains("document 2 produced no text"),
"error should name document 2: {error}"
);
}
#[derive(Clone, Default)]
struct OneAtATimeReversedLatency;
impl Transport<Local<crate::operation::Embedding>> for OneAtATimeReversedLatency {
fn send(
&self,
documents: Vec<String>,
_exchange: Exchange,
) -> Opening<Step<crate::operation::Embedding>> {
let Some(position) = documents
.first()
.and_then(|text| text.rsplit_once('t'))
.and_then(|(_, n)| n.parse::<u64>().ok())
else {
return Opening::failed(ProviderError::Provider(format!(
"could not read a text position out of {documents:?}; \
this mock cannot invert completion order without it"
)));
};
let delay = 60u64.saturating_sub(position * 10);
Opening::new(async move {
tokio::time::sleep(std::time::Duration::from_millis(delay)).await;
Ok(Opened::new(futures::stream::iter([Ok(Step::End(
MockEmbeddings::embed(documents),
))])))
})
}
}
#[tokio::test]
async fn test_build_order_survives_one_text_per_batch_finishing_backwards() {
let doc = NTexts::new(0, 6);
let expected = doc.expected();
let result = EmbeddingsBuilder::new(Model::new(batches(1), OneAtATimeReversedLatency))
.document(doc)
.unwrap()
.build()
.await
.unwrap();
assert_eq!(result.len(), 1);
assert_eq!(returned(&result[0].1), expected);
}