Skip to main content

rig_core/test_utils/
embeddings.rs

1//! Embedding helpers for deterministic tests.
2
3use crate::driver::{Exchange, Local, Model, Opened, Opening, Step, Transport};
4use crate::error::ProviderError;
5use crate::wire::Capabilities;
6use crate::{
7    Embed,
8    embeddings::{
9        Embedding, EmbeddingResponse,
10        embed::{EmbedError, TextEmbedder},
11    },
12};
13
14/// The mock embedding runtime: every text embeds to one fixed vector, five
15/// texts to a request at ten dimensions. It is the transport of a [`Local`]
16/// embedding wire ([`Self::model`]).
17#[derive(Clone, Copy, Debug, Default, PartialEq)]
18pub struct MockEmbeddings;
19
20/// A deterministic embedding model that returns a fixed vector for each input document.
21pub type MockEmbeddingModel = Model<Local<crate::operation::Embedding>, MockEmbeddings>;
22
23impl MockEmbeddings {
24    /// The mock embedding model: its local wire over this runtime.
25    pub fn model() -> MockEmbeddingModel {
26        Model::new(
27            Local::new(super::MOCK_PROVIDER).with_capabilities(Capabilities::embedding(5, 10)),
28            Self,
29        )
30    }
31}
32
33impl Transport<Local<crate::operation::Embedding>> for MockEmbeddings {
34    fn send(
35        &self,
36        texts: Vec<String>,
37        _exchange: Exchange,
38    ) -> Opening<Step<crate::operation::Embedding>> {
39        let response = Self::embed(texts);
40        Opening::ready(Opened::new(futures::stream::iter([
41            Ok::<_, ProviderError>(Step::End(response)),
42        ])))
43    }
44}
45
46impl MockEmbeddings {
47    /// The reply this runtime gives for `texts`: one fixed ten-dimension
48    /// vector per text, in order.
49    pub fn embed(texts: Vec<String>) -> EmbeddingResponse {
50        EmbeddingResponse {
51            provider: super::MOCK_PROVIDER.to_owned(),
52            ..EmbeddingResponse::new(
53                texts
54                    .into_iter()
55                    .map(|document| Embedding {
56                        document,
57                        vec: vec![0.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9],
58                    })
59                    .collect(),
60            )
61        }
62    }
63}
64
65/// A test document that contributes one text fragment to an embedding request.
66#[derive(Clone, Debug)]
67pub struct MockTextDocument {
68    /// Stable document identifier used by tests.
69    pub id: String,
70    /// Text to embed.
71    pub text: String,
72}
73
74impl MockTextDocument {
75    /// Create a single-text embedding fixture.
76    pub fn new(id: impl Into<String>, text: impl Into<String>) -> Self {
77        Self {
78            id: id.into(),
79            text: text.into(),
80        }
81    }
82}
83
84impl Embed for MockTextDocument {
85    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
86        embedder.embed(self.text.clone());
87        Ok(())
88    }
89}
90
91/// A test document that contributes multiple text fragments to an embedding request.
92#[derive(Clone, Debug)]
93pub struct MockMultiTextDocument {
94    /// Stable document identifier used by tests.
95    pub id: String,
96    /// Text fragments to embed.
97    pub texts: Vec<String>,
98}
99
100impl MockMultiTextDocument {
101    /// Create a multi-text embedding fixture.
102    pub fn new(id: impl Into<String>, texts: impl IntoIterator<Item = impl Into<String>>) -> Self {
103        Self {
104            id: id.into(),
105            texts: texts.into_iter().map(Into::into).collect(),
106        }
107    }
108}
109
110impl Embed for MockMultiTextDocument {
111    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
112        for text in &self.texts {
113            embedder.embed(text.clone());
114        }
115        Ok(())
116    }
117}