use async_trait::async_trait;
use rskit_ai::Usage;
use rskit_ai::semconv;
use rskit_component::{Component, Health};
use rskit_errors::AppResult;
use rskit_observability::set_span_attribute;
use crate::{EmbedInput, EmbedRequest, EmbedResponse, Embedding, Provider};
#[derive(Debug, Clone)]
pub struct InMemoryProvider {
dimensions: usize,
}
impl InMemoryProvider {
#[must_use]
pub const fn new(dimensions: usize) -> Self {
Self { dimensions }
}
fn vector_for(&self, input: &EmbedInput) -> Vec<f32> {
let bytes: Vec<u8> = match input {
EmbedInput::Text(text) => text.as_bytes().to_vec(),
EmbedInput::Image(asset) | EmbedInput::Audio(asset) | EmbedInput::Video(asset) => {
serde_json::to_vec(asset).unwrap_or_default()
}
};
(0..self.dimensions)
.map(|idx| {
let seed = u32::try_from(idx).unwrap_or(u32::MAX);
let sum = bytes.iter().enumerate().fold(seed, |acc, (pos, byte)| {
let factor = u32::try_from(pos + idx + 1).unwrap_or(u32::MAX);
acc.wrapping_add(u32::from(*byte) * factor)
});
f32::from(u16::try_from(sum % 1000).unwrap_or(0)) / 1000.0
})
.collect()
}
}
impl Default for InMemoryProvider {
fn default() -> Self {
Self::new(8)
}
}
#[async_trait]
impl Provider for InMemoryProvider {
async fn embed(&self, req: EmbedRequest) -> AppResult<EmbedResponse> {
let span = tracing::info_span!(
"embedding.embed",
"gen_ai.system" = "in_memory",
"gen_ai.operation.name" = semconv::Operation::Embedding.as_str(),
"gen_ai.request.model" = req.model.name.as_str(),
"embedding.input_count" = req.inputs.len(),
);
set_span_attribute(&span, semconv::SYSTEM, "in_memory");
set_span_attribute(
&span,
semconv::OPERATION_NAME,
semconv::Operation::Embedding.as_str(),
);
set_span_attribute(&span, semconv::REQUEST_MODEL, req.model.name.as_str());
let _span = span.entered();
let embeddings = req
.inputs
.iter()
.enumerate()
.map(|(index, input)| Embedding::new(self.vector_for(input), index))
.collect();
Ok(EmbedResponse {
embeddings,
model: req.model,
usage: Usage::default(),
})
}
async fn embed_batch(&self, reqs: Vec<EmbedRequest>) -> AppResult<Vec<EmbedResponse>> {
let mut responses = Vec::with_capacity(reqs.len());
for req in reqs {
responses.push(self.embed(req).await?);
}
Ok(responses)
}
}
impl rskit_provider::Provider for InMemoryProvider {
fn name(&self) -> &'static str {
"in_memory_embedding"
}
}
#[async_trait]
impl rskit_provider::RequestResponse<EmbedRequest, EmbedResponse> for InMemoryProvider {
async fn execute(&self, input: EmbedRequest) -> AppResult<EmbedResponse> {
self.embed(input).await
}
}
#[async_trait]
impl Component for InMemoryProvider {
fn name(&self) -> &'static str {
"rskit-embedding.in_memory"
}
async fn start(&self) -> AppResult<()> {
Ok(())
}
async fn stop(&self) -> AppResult<()> {
Ok(())
}
fn health(&self) -> Health {
Health::healthy(self.name())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::EmbeddingOptions;
use rskit_ai::{Capabilities, Model, Provider as ModelProvider};
fn model() -> Model {
Model {
name: "embed-test".into(),
provider: ModelProvider::Custom("memory".into()),
version: None,
capabilities: Capabilities::default(),
}
}
#[tokio::test]
async fn deterministic_adapter_embeds_inputs() {
let provider = InMemoryProvider::new(4);
let req = EmbedRequest {
model: model(),
inputs: vec![
EmbedInput::Text("hello".into()),
EmbedInput::Text("world".into()),
],
options: EmbeddingOptions::default(),
};
let response = provider.embed(req.clone()).await.expect("embed");
let again = provider.embed(req).await.expect("embed again");
assert_eq!(response.embeddings, again.embeddings);
assert_eq!(response.embeddings[0].dimensions, 4);
assert_eq!(response.embeddings[1].index, 1);
assert_eq!(response.usage, Usage::default());
}
#[tokio::test]
async fn batch_returns_one_response_per_request() {
let provider = InMemoryProvider::default();
let req = EmbedRequest {
model: model(),
inputs: vec![EmbedInput::Text("x".into())],
options: EmbeddingOptions::default(),
};
let responses = provider
.embed_batch(vec![req.clone(), req])
.await
.expect("batch");
assert_eq!(responses.len(), 2);
assert_eq!(responses[0].embeddings[0].dimensions, 8);
}
}