rskit-embedding 0.2.0-alpha.2

Embedding provider abstractions for vector search
Documentation
//! Deterministic in-memory embedding adapter for tests.

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

/// Deterministic embedding provider for tests and examples.
#[derive(Debug, Clone)]
pub struct InMemoryProvider {
    dimensions: usize,
}

impl InMemoryProvider {
    /// Create a deterministic provider with fixed vector dimensions.
    #[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);
    }
}