use rig::{
Embed, OneOrMany,
embeddings::{Embedding, EmbeddingModel},
vector_store::{
InsertDocuments, TopNResults, VectorStoreError, VectorStoreIndex, VectorStoreIndexDyn,
request::{Filter, FilterError, SearchFilter, VectorSearchRequest},
},
wasm_compat::WasmBoxedFuture,
};
use scylla::{
client::{Compression, session::Session, session_builder::SessionBuilder},
statement::prepared::PreparedStatement,
value::CqlValue,
};
use serde::{Deserialize, Serialize};
use std::{
collections::HashMap,
hash::{DefaultHasher, Hash, Hasher},
sync::{Arc, RwLock},
};
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,
cache: Arc<RwLock<HashMap<u64, PreparedStatement>>>,
}
fn cql_value_from_json(value: serde_json::Value) -> Result<CqlValue, FilterError> {
use scylla::value::CqlVarint;
use serde_json::Value;
match value {
Value::Bool(b) => Ok(CqlValue::Boolean(b)),
Value::Number(n) => {
if let Some(i) = n.as_i64() {
Ok(CqlValue::BigInt(i))
} else if let Some(u) = n.as_u64() {
let mut bytes = vec![0u8];
bytes.extend_from_slice(&u.to_be_bytes());
Ok(CqlValue::Varint(CqlVarint::from_signed_bytes_be(bytes)))
} else if let Some(f) = n.as_f64() {
Ok(CqlValue::Double(f))
} else {
Err(FilterError::Expected {
expected: "Valid number".into(),
got: "Invalid number".into(),
})
}
}
Value::String(s) => Ok(CqlValue::Text(s)),
Value::Array(arr) => Ok(CqlValue::List(
arr.into_iter()
.map(cql_value_from_json)
.collect::<Result<_, _>>()?,
)),
Value::Object(map) => {
let pairs = map
.into_iter()
.map(|(k, v)| Ok((CqlValue::Text(k), cql_value_from_json(v)?)))
.collect::<Result<Vec<_>, FilterError>>()?;
Ok(CqlValue::Map(pairs))
}
Value::Null => Ok(CqlValue::Empty),
}
}
#[derive(Clone, Debug)]
pub struct ScyllaSearchFilter {
condition: String,
params: Vec<CqlValue>,
}
impl std::hash::Hash for ScyllaSearchFilter {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.condition.hash(state)
}
}
impl SearchFilter for ScyllaSearchFilter {
type Value = CqlValue;
fn eq(key: impl AsRef<str>, value: Self::Value) -> Self {
Self {
condition: format!("{} = ?", key.as_ref()),
params: vec![value],
}
}
fn gt(key: impl AsRef<str>, value: Self::Value) -> Self {
Self {
condition: format!("{} > ?", key.as_ref()),
params: vec![value],
}
}
fn lt(key: impl AsRef<str>, value: Self::Value) -> Self {
Self {
condition: format!("{} < ?", key.as_ref()),
params: vec![value],
}
}
fn and(self, rhs: Self) -> Self {
Self {
condition: format!("({}) AND ({})", self.condition, rhs.condition),
params: self.params.into_iter().chain(rhs.params).collect(),
}
}
fn or(self, rhs: Self) -> Self {
Self {
condition: format!("({}) OR ({})", self.condition, rhs.condition),
params: self.params.into_iter().chain(rhs.params).collect(),
}
}
}
impl ScyllaSearchFilter {
fn params(&self) -> &[CqlValue] {
self.params.as_slice()
}
#[allow(clippy::should_implement_trait)]
pub fn not(self) -> Self {
Self {
condition: format!("NOT ({})", self.condition),
..self
}
}
pub fn gte(key: String, value: <Self as SearchFilter>::Value) -> Self {
Self {
condition: format!("{key} >= ?"),
params: vec![value],
}
}
pub fn lte(key: String, value: <Self as SearchFilter>::Value) -> Self {
Self {
condition: format!("{key} <= ?"),
params: vec![value],
}
}
pub fn ne(key: String, value: <Self as SearchFilter>::Value) -> Self {
Self {
condition: format!("{key} != ?"),
params: vec![value],
}
}
pub fn member(key: String, values: Vec<<Self as SearchFilter>::Value>) -> Self {
let placeholders = vec!["?"; values.len()].join(", ");
Self {
condition: format!("{key} IN ({placeholders})"),
params: values,
}
}
}
impl TryFrom<Filter<serde_json::Value>> for ScyllaSearchFilter {
type Error = FilterError;
fn try_from(value: Filter<serde_json::Value>) -> Result<Self, Self::Error> {
match value {
Filter::Eq(k, val) => Ok(ScyllaSearchFilter::eq(k, cql_value_from_json(val)?)),
Filter::Gt(k, val) => Ok(ScyllaSearchFilter::gt(k, cql_value_from_json(val)?)),
Filter::Lt(k, val) => Ok(ScyllaSearchFilter::lt(k, cql_value_from_json(val)?)),
Filter::And(l, r) => Ok(Self::try_from(*l)?.and(Self::try_from(*r)?)),
Filter::Or(l, r) => Ok(Self::try_from(*l)?.or(Self::try_from(*r)?)),
}
}
}
impl<M> ScyllaDbVectorStore<M>
where
M: EmbeddingModel,
{
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,
cache: Default::default(),
})
}
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())
}
async fn get_filter_statement_or_default(
&self,
req: &VectorSearchRequest<ScyllaSearchFilter>,
) -> Result<PreparedStatement, VectorStoreError> {
if let Some(filter) = req.filter() {
let mut hasher = DefaultHasher::new();
filter.hash(&mut hasher);
let filter_hash = hasher.finish();
let statement = if let Some(cached) = self
.cache
.read()
.ok()
.and_then(|cache| cache.get(&filter_hash).cloned())
{
cached
} else {
let query = format!(
"SELECT id, vector, metadata, created_at FROM {}.{} WHERE {} ALLOW FILTERING",
self.keyspace, self.table, filter.condition
);
let prepared = self
.session
.prepare(query)
.await
.map_err(|e| VectorStoreError::DatastoreError(e.into()))?;
let mut cache = self.cache.write().map_err(|e| {
VectorStoreError::DatastoreError(
format!("Error writing statement cache: {e}").into(),
)
})?;
cache.insert(filter_hash, prepared.clone());
prepared
};
Ok(statement)
} else {
Ok(self.search_stmt.clone())
}
}
}
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> VectorStoreIndex for ScyllaDbVectorStore<M>
where
M: EmbeddingModel + std::marker::Sync + Send,
{
type Filter = ScyllaSearchFilter;
async fn top_n<T: for<'a> Deserialize<'a> + Send>(
&self,
req: VectorSearchRequest<ScyllaSearchFilter>,
) -> Result<Vec<(f64, String, T)>, VectorStoreError> {
let query_vector = self.generate_query_vector(req.query()).await?;
let statement = self.get_filter_statement_or_default(&req).await?;
let params = req
.filter()
.as_ref()
.map(ScyllaSearchFilter::params)
.unwrap_or([].as_slice());
let results = self
.session
.execute_unpaged(&statement, params)
.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;
if req.threshold().is_some_and(|threshold| score < threshold) {
continue;
}
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(req.samples() as usize);
Ok(candidates)
}
async fn top_n_ids(
&self,
req: VectorSearchRequest<ScyllaSearchFilter>,
) -> Result<Vec<(f64, String)>, VectorStoreError> {
let query_vector = self.generate_query_vector(req.query()).await?;
let statement = self.get_filter_statement_or_default(&req).await?;
let params = req
.filter()
.as_ref()
.map(ScyllaSearchFilter::params)
.unwrap_or_default();
let results = self
.session
.execute_unpaged(&statement, params)
.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;
if req.threshold().is_some_and(|threshold| score < threshold) {
continue;
}
candidates.push((score, id.to_string()));
}
candidates.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap());
candidates.truncate(req.samples() as usize);
Ok(candidates)
}
}
impl<M> VectorStoreIndexDyn for ScyllaDbVectorStore<M>
where
M: EmbeddingModel + Sync + Send,
{
fn top_n<'a>(
&'a self,
req: VectorSearchRequest<Filter<serde_json::Value>>,
) -> WasmBoxedFuture<'a, TopNResults> {
Box::pin(async move {
let req = req.try_map_filter(ScyllaSearchFilter::try_from)?;
let results = <Self as VectorStoreIndex>::top_n::<serde_json::Value>(self, req).await?;
Ok(results)
})
}
fn top_n_ids<'a>(
&'a self,
req: VectorSearchRequest<Filter<serde_json::Value>>,
) -> WasmBoxedFuture<'a, Result<Vec<(f64, String)>, VectorStoreError>> {
Box::pin(async move {
let req = req.try_map_filter(ScyllaSearchFilter::try_from)?;
let results = <Self as VectorStoreIndex>::top_n_ids(self, req).await?;
Ok(results)
})
}
}
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)))
}