Skip to main content

rskit_embedding/
in_memory.rs

1//! Deterministic in-memory embedding adapter for tests.
2
3use async_trait::async_trait;
4use rskit_ai::Usage;
5use rskit_ai::semconv;
6use rskit_component::{Component, Health};
7use rskit_errors::AppResult;
8use rskit_observability::set_span_attribute;
9
10use crate::{EmbedInput, EmbedRequest, EmbedResponse, Embedding, Provider};
11
12/// Deterministic embedding provider for tests and examples.
13#[derive(Debug, Clone)]
14pub struct InMemoryProvider {
15    dimensions: usize,
16}
17
18impl InMemoryProvider {
19    /// Create a deterministic provider with fixed vector dimensions.
20    #[must_use]
21    pub const fn new(dimensions: usize) -> Self {
22        Self { dimensions }
23    }
24
25    fn vector_for(&self, input: &EmbedInput) -> Vec<f32> {
26        let bytes: Vec<u8> = match input {
27            EmbedInput::Text(text) => text.as_bytes().to_vec(),
28            EmbedInput::Image(asset) | EmbedInput::Audio(asset) | EmbedInput::Video(asset) => {
29                serde_json::to_vec(asset).unwrap_or_default()
30            }
31        };
32        (0..self.dimensions)
33            .map(|idx| {
34                let seed = u32::try_from(idx).unwrap_or(u32::MAX);
35                let sum = bytes.iter().enumerate().fold(seed, |acc, (pos, byte)| {
36                    let factor = u32::try_from(pos + idx + 1).unwrap_or(u32::MAX);
37                    acc.wrapping_add(u32::from(*byte) * factor)
38                });
39                f32::from(u16::try_from(sum % 1000).unwrap_or(0)) / 1000.0
40            })
41            .collect()
42    }
43}
44
45impl Default for InMemoryProvider {
46    fn default() -> Self {
47        Self::new(8)
48    }
49}
50
51#[async_trait]
52impl Provider for InMemoryProvider {
53    async fn embed(&self, req: EmbedRequest) -> AppResult<EmbedResponse> {
54        let span = tracing::info_span!(
55            "embedding.embed",
56            "gen_ai.system" = "in_memory",
57            "gen_ai.operation.name" = semconv::Operation::Embedding.as_str(),
58            "gen_ai.request.model" = req.model.name.as_str(),
59            "embedding.input_count" = req.inputs.len(),
60        );
61        set_span_attribute(&span, semconv::SYSTEM, "in_memory");
62        set_span_attribute(
63            &span,
64            semconv::OPERATION_NAME,
65            semconv::Operation::Embedding.as_str(),
66        );
67        set_span_attribute(&span, semconv::REQUEST_MODEL, req.model.name.as_str());
68        let _span = span.entered();
69        let embeddings = req
70            .inputs
71            .iter()
72            .enumerate()
73            .map(|(index, input)| Embedding::new(self.vector_for(input), index))
74            .collect();
75        Ok(EmbedResponse {
76            embeddings,
77            model: req.model,
78            usage: Usage::default(),
79        })
80    }
81
82    async fn embed_batch(&self, reqs: Vec<EmbedRequest>) -> AppResult<Vec<EmbedResponse>> {
83        let mut responses = Vec::with_capacity(reqs.len());
84        for req in reqs {
85            responses.push(self.embed(req).await?);
86        }
87        Ok(responses)
88    }
89}
90
91impl rskit_provider::Provider for InMemoryProvider {
92    fn name(&self) -> &'static str {
93        "in_memory_embedding"
94    }
95}
96
97#[async_trait]
98impl rskit_provider::RequestResponse<EmbedRequest, EmbedResponse> for InMemoryProvider {
99    async fn execute(&self, input: EmbedRequest) -> AppResult<EmbedResponse> {
100        self.embed(input).await
101    }
102}
103
104#[async_trait]
105impl Component for InMemoryProvider {
106    fn name(&self) -> &'static str {
107        "rskit-embedding.in_memory"
108    }
109
110    async fn start(&self) -> AppResult<()> {
111        Ok(())
112    }
113
114    async fn stop(&self) -> AppResult<()> {
115        Ok(())
116    }
117
118    fn health(&self) -> Health {
119        Health::healthy(self.name())
120    }
121}
122
123#[cfg(test)]
124mod tests {
125    use super::*;
126    use crate::EmbeddingOptions;
127    use rskit_ai::{Capabilities, Model, Provider as ModelProvider};
128
129    fn model() -> Model {
130        Model {
131            name: "embed-test".into(),
132            provider: ModelProvider::Custom("memory".into()),
133            version: None,
134            capabilities: Capabilities::default(),
135        }
136    }
137
138    #[tokio::test]
139    async fn deterministic_adapter_embeds_inputs() {
140        let provider = InMemoryProvider::new(4);
141        let req = EmbedRequest {
142            model: model(),
143            inputs: vec![
144                EmbedInput::Text("hello".into()),
145                EmbedInput::Text("world".into()),
146            ],
147            options: EmbeddingOptions::default(),
148        };
149        let response = provider.embed(req.clone()).await.expect("embed");
150        let again = provider.embed(req).await.expect("embed again");
151        assert_eq!(response.embeddings, again.embeddings);
152        assert_eq!(response.embeddings[0].dimensions, 4);
153        assert_eq!(response.embeddings[1].index, 1);
154        assert_eq!(response.usage, Usage::default());
155    }
156
157    #[tokio::test]
158    async fn batch_returns_one_response_per_request() {
159        let provider = InMemoryProvider::default();
160        let req = EmbedRequest {
161            model: model(),
162            inputs: vec![EmbedInput::Text("x".into())],
163            options: EmbeddingOptions::default(),
164        };
165        let responses = provider
166            .embed_batch(vec![req.clone(), req])
167            .await
168            .expect("batch");
169        assert_eq!(responses.len(), 2);
170        assert_eq!(responses[0].embeddings[0].dimensions, 8);
171    }
172}