rskit_embedding/
in_memory.rs1use 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#[derive(Debug, Clone)]
14pub struct InMemoryProvider {
15 dimensions: usize,
16}
17
18impl InMemoryProvider {
19 #[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}