rig_core/test_utils/
embeddings.rs1use 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#[derive(Clone, Copy, Debug, Default, PartialEq)]
18pub struct MockEmbeddings;
19
20pub type MockEmbeddingModel = Model<Local<crate::operation::Embedding>, MockEmbeddings>;
22
23impl MockEmbeddings {
24 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 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#[derive(Clone, Debug)]
67pub struct MockTextDocument {
68 pub id: String,
70 pub text: String,
72}
73
74impl MockTextDocument {
75 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#[derive(Clone, Debug)]
93pub struct MockMultiTextDocument {
94 pub id: String,
96 pub texts: Vec<String>,
98}
99
100impl MockMultiTextDocument {
101 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}