Skip to main content

rig_vectorize/
lib.rs

1//! Cloudflare Vectorize integration for the Rig framework.
2//!
3//! This crate provides a vector store implementation using Cloudflare Vectorize,
4//! a globally distributed vector database built for AI applications.
5//!
6//! # Example
7//!
8//! ```ignore
9//! use rig_core::client::ProviderClient;
10//! use rig_core::providers::openai;
11//! use rig_vectorize::VectorizeVectorStore;
12//!
13//! let openai = openai::Client::from_env()?;
14//! let embedding_model = openai.embedding_model(openai::TEXT_EMBEDDING_3_SMALL);
15//!
16//! let vector_store = VectorizeVectorStore::new(
17//!     embedding_model,
18//!     "your-account-id",
19//!     "your-index-name",
20//!     std::env::var("CLOUDFLARE_API_TOKEN").unwrap(),
21//! );
22//! ```
23
24mod client;
25
26// Re-export client types
27pub use client::{
28    DeleteByIdsRequest, DeleteResult, ListVectorsResult, QueryRequest, QueryResult, ReturnMetadata,
29    UpsertRequest, UpsertResult, VectorIdEntry, VectorInput, VectorMatch, VectorizeClient,
30    VectorizeError, VectorizeFilter,
31};
32
33use client::{QueryRequest as ApiQueryRequest, VectorInput as ApiVectorInput};
34use rig_core::embeddings::EmbeddingModel;
35use rig_core::vector_store::request::VectorSearchRequest;
36use rig_core::vector_store::{InsertDocuments, VectorStoreError, VectorStoreIndex};
37use rig_core::{Embed, embeddings::Embedding};
38use serde::{Deserialize, Serialize};
39use uuid::Uuid;
40
41impl From<VectorizeError> for VectorStoreError {
42    fn from(err: VectorizeError) -> Self {
43        VectorStoreError::datastore(err)
44    }
45}
46
47/// A vector store backed by Cloudflare Vectorize.
48///
49/// This struct implements [`VectorStoreIndex`] to provide vector similarity search
50/// using Cloudflare's globally distributed Vectorize service.
51#[derive(Debug, Clone)]
52pub struct VectorizeVectorStore<M> {
53    /// The embedding model used to generate query embeddings.
54    model: M,
55    /// The HTTP client for Vectorize API.
56    client: VectorizeClient,
57}
58
59impl<M> VectorizeVectorStore<M> {
60    /// Creates a new Vectorize vector store.
61    ///
62    /// # Arguments
63    /// * `model` - The embedding model to use for query embedding
64    /// * `account_id` - Cloudflare account ID
65    /// * `index_name` - Name of the Vectorize index
66    /// * `api_token` - Cloudflare API token with Vectorize read permissions
67    pub fn new(
68        model: M,
69        account_id: impl Into<String>,
70        index_name: impl Into<String>,
71        api_token: impl Into<String>,
72    ) -> Self {
73        Self {
74            model,
75            client: VectorizeClient::new(account_id, index_name, api_token),
76        }
77    }
78}
79
80impl<M> VectorizeVectorStore<M>
81where
82    M: EmbeddingModel + Sync + Send,
83{
84    /// Validates the filter, embeds the query, and returns the threshold-filtered matches.
85    async fn query_matches(
86        &self,
87        req: &VectorSearchRequest<VectorizeFilter>,
88        return_metadata: ReturnMetadata,
89    ) -> Result<Vec<VectorMatch>, VectorStoreError> {
90        if let Some(filter) = req.filter() {
91            filter.validate()?;
92        }
93
94        let embedding = self.model.embed_text(req.query()).await?;
95
96        let query_request = ApiQueryRequest {
97            vector: embedding.vec,
98            top_k: req.samples(),
99            return_values: Some(false),
100            return_metadata: Some(return_metadata),
101            filter: req.filter().as_ref().map(|f| f.clone().into_inner()),
102        };
103
104        let result = self.client.query(query_request).await?;
105
106        Ok(result
107            .matches
108            .into_iter()
109            .filter(|m| req.threshold().is_none_or(|t| m.score >= t))
110            .collect())
111    }
112}
113
114impl<M> VectorStoreIndex for VectorizeVectorStore<M>
115where
116    M: EmbeddingModel + Sync + Send,
117{
118    type Filter = VectorizeFilter;
119
120    async fn top_n<T: for<'a> Deserialize<'a> + Send>(
121        &self,
122        req: VectorSearchRequest<Self::Filter>,
123    ) -> Result<Vec<(f64, String, T)>, VectorStoreError> {
124        let matches = self.query_matches(&req, ReturnMetadata::All).await?;
125
126        // Convert results to the expected format
127        let results = matches
128            .into_iter()
129            .map(|m| {
130                let metadata = m.metadata.unwrap_or(serde_json::Value::Null);
131                let doc: T = serde_json::from_value(metadata)?;
132                Ok((m.score, m.id, doc))
133            })
134            .collect::<Result<Vec<_>, serde_json::Error>>()?;
135
136        Ok(results)
137    }
138
139    async fn top_n_ids(
140        &self,
141        req: VectorSearchRequest<Self::Filter>,
142    ) -> Result<Vec<(f64, String)>, VectorStoreError> {
143        let matches = self.query_matches(&req, ReturnMetadata::None).await?;
144
145        // Convert results to (score, id) tuples
146        Ok(matches.into_iter().map(|m| (m.score, m.id)).collect())
147    }
148}
149
150impl<M> InsertDocuments for VectorizeVectorStore<M>
151where
152    M: EmbeddingModel + Sync + Send,
153{
154    async fn insert_documents<Doc: Serialize + Embed + Send>(
155        &self,
156        documents: Vec<(Doc, Vec<Embedding>)>,
157    ) -> Result<(), VectorStoreError> {
158        let vectors =
159            rig_core::vector_store::flatten_embedded(documents, |metadata, embedding| {
160                Ok(ApiVectorInput {
161                    id: Uuid::new_v4().to_string(),
162                    values: embedding.vec,
163                    metadata: Some(metadata.clone()),
164                    namespace: None,
165                })
166            })?;
167
168        if vectors.is_empty() {
169            return Ok(());
170        }
171
172        tracing::debug!("Upserting {} vectors to Vectorize", vectors.len());
173
174        const BATCH_SIZE: usize = 1000;
175
176        for batch in vectors.chunks(BATCH_SIZE) {
177            let request = UpsertRequest {
178                vectors: batch.to_vec(),
179            };
180
181            self.client.upsert(request).await?;
182        }
183
184        Ok(())
185    }
186}