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 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}