1mod client;
25
26pub 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#[derive(Debug, Clone)]
52pub struct VectorizeVectorStore<M> {
53 model: M,
55 client: VectorizeClient,
57}
58
59impl<M> VectorizeVectorStore<M> {
60 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 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 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 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}