use rig::{
Embed, OneOrMany,
embeddings::{Embedding, EmbeddingModel},
vector_store::{InsertDocuments, VectorStoreError, VectorStoreIndex},
};
use scylla::{
client::{Compression, session::Session, session_builder::SessionBuilder},
statement::prepared::PreparedStatement,
};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use uuid::Uuid;
pub struct ScyllaDbVectorStore<M: EmbeddingModel> {
model: M,
pub session: Arc<Session>,
keyspace: String,
table: String,
dimensions: usize,
insert_stmt: PreparedStatement,
search_stmt: PreparedStatement,
get_by_id_stmt: PreparedStatement,
}
#[derive(Debug, Serialize, Deserialize)]
struct VectorRecord {
id: Uuid,
vector: Vec<f32>,
metadata: String, created_at: i64, }
impl<M: EmbeddingModel> ScyllaDbVectorStore<M> {
pub async fn new(
model: M,
session: Session,
keyspace: &str,
table: &str,
dimensions: usize,
) -> Result<Self, VectorStoreError> {
let session = Arc::new(session);
let create_keyspace_cql = format!(
"CREATE KEYSPACE IF NOT EXISTS {keyspace} WITH REPLICATION = {{
'class': 'SimpleStrategy',
'replication_factor': 1
}}"
);
session
.query_unpaged(create_keyspace_cql, &[])
.await
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let create_table_cql = format!(
"CREATE TABLE IF NOT EXISTS {keyspace}.{table} (
id UUID PRIMARY KEY,
vector LIST<FLOAT>,
metadata TEXT,
created_at BIGINT
)"
);
session
.query_unpaged(create_table_cql, &[])
.await
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let insert_stmt = session
.prepare(format!(
"INSERT INTO {keyspace}.{table} (id, vector, metadata, created_at) VALUES (?, ?, ?, ?)"
))
.await
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let search_stmt = session
.prepare(format!(
"SELECT id, vector, metadata, created_at FROM {keyspace}.{table}"
))
.await
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let get_by_id_stmt = session
.prepare(format!(
"SELECT id, vector, metadata, created_at FROM {keyspace}.{table} WHERE id = ?"
))
.await
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
Ok(Self {
model,
session,
keyspace: keyspace.to_string(),
table: table.to_string(),
dimensions,
insert_stmt,
search_stmt,
get_by_id_stmt,
})
}
pub fn session(&self) -> &Arc<Session> {
&self.session
}
pub fn keyspace(&self) -> &str {
&self.keyspace
}
pub fn table(&self) -> &str {
&self.table
}
pub async fn get_by_id<T: for<'a> Deserialize<'a> + Send>(
&self,
id: &str,
) -> Result<Option<T>, VectorStoreError> {
let uuid =
Uuid::parse_str(id).map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let result = self
.session
.execute_unpaged(&self.get_by_id_stmt, (uuid,))
.await
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let rows_result = result
.into_rows_result()
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
if let Some(first_row) = rows_result
.rows::<(Uuid, Vec<f32>, String, i64)>()
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?
.next()
{
let (_, _, metadata, _) =
first_row.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let payload: T = serde_json::from_str(&metadata)?;
return Ok(Some(payload));
}
Ok(None)
}
fn cosine_similarity(vec1: &[f32], vec2: &[f32]) -> f32 {
let dot_product: f32 = vec1.iter().zip(vec2.iter()).map(|(a, b)| a * b).sum();
let norm1: f32 = vec1.iter().map(|x| x * x).sum::<f32>().sqrt();
let norm2: f32 = vec2.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm1 == 0.0 || norm2 == 0.0 {
0.0
} else {
dot_product / (norm1 * norm2)
}
}
async fn generate_query_vector(&self, query: &str) -> Result<Vec<f32>, VectorStoreError> {
let embedding = self.model.embed_text(query).await?;
Ok(embedding.vec.iter().map(|&x| x as f32).collect())
}
}
impl<Model> InsertDocuments for ScyllaDbVectorStore<Model>
where
Model: EmbeddingModel + Send + Sync,
{
async fn insert_documents<Doc: Serialize + Embed + Send>(
&self,
documents: Vec<(Doc, OneOrMany<Embedding>)>,
) -> Result<(), VectorStoreError> {
for (document, embeddings) in documents {
let metadata = serde_json::to_string(&document)?;
let now = chrono::Utc::now().timestamp();
for embedding in embeddings.into_iter() {
let vector: Vec<f32> = embedding.vec.into_iter().map(|x| x as f32).collect();
if vector.len() != self.dimensions {
return Err(VectorStoreError::DatastoreError(
format!(
"Vector dimension mismatch: expected {}, got {}",
self.dimensions,
vector.len()
)
.into(),
));
}
let id = Uuid::new_v4();
self.session
.execute_unpaged(&self.insert_stmt, (id, vector, &metadata, now))
.await
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
}
}
Ok(())
}
}
impl<M: EmbeddingModel + std::marker::Sync + Send> VectorStoreIndex for ScyllaDbVectorStore<M> {
async fn top_n<T: for<'a> Deserialize<'a> + Send>(
&self,
query: &str,
n: usize,
) -> Result<Vec<(f64, String, T)>, VectorStoreError> {
let query_vector = self.generate_query_vector(query).await?;
let results = self
.session
.execute_unpaged(&self.search_stmt, &[])
.await
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let rows_result = results
.into_rows_result()
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let mut candidates = Vec::new();
for row_result in rows_result
.rows::<(Uuid, Vec<f32>, String, i64)>()
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?
{
let (id, vector, metadata, _) =
row_result.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let similarity = Self::cosine_similarity(&query_vector, &vector);
let score = similarity as f64;
let payload: T = serde_json::from_str(&metadata)?;
candidates.push((score, id.to_string(), payload));
}
candidates.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap());
candidates.truncate(n);
Ok(candidates)
}
async fn top_n_ids(
&self,
query: &str,
n: usize,
) -> Result<Vec<(f64, String)>, VectorStoreError> {
let query_vector = self.generate_query_vector(query).await?;
let results = self
.session
.execute_unpaged(&self.search_stmt, &[])
.await
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let rows_result = results
.into_rows_result()
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let mut candidates = Vec::new();
for row_result in rows_result
.rows::<(Uuid, Vec<f32>, String, i64)>()
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?
{
let (id, vector, _, _) =
row_result.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))?;
let similarity = Self::cosine_similarity(&query_vector, &vector);
let score = similarity as f64;
candidates.push((score, id.to_string()));
}
candidates.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap());
candidates.truncate(n);
Ok(candidates)
}
}
pub async fn create_session(uri: &str) -> Result<Session, VectorStoreError> {
SessionBuilder::new()
.known_node(uri)
.compression(Some(Compression::Lz4))
.build()
.await
.map_err(|e| VectorStoreError::DatastoreError(Box::new(e)))
}