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