1pub mod vector_index;
54use std::str::FromStr;
55
56use futures::TryStreamExt;
57use neo4rs::*;
58use rig_core::vector_store::{VectorStoreError, request::SearchFilter};
59use serde::{Deserialize, Serialize};
60use vector_index::{IndexConfig, Neo4jVectorIndex, VectorSimilarityFunction};
61
62pub struct Neo4jClient {
63 pub graph: Graph,
64}
65
66#[derive(Clone, Debug, Serialize, Deserialize)]
71pub struct Neo4jSearchFilter(String);
72
73impl SearchFilter for Neo4jSearchFilter {
74 type Value = serde_json::Value;
75
76 fn eq(key: impl AsRef<str>, value: Self::Value) -> Self {
77 Self(format!("n.{} = {}", key.as_ref(), serialize_cypher(value)))
78 }
79
80 fn gt(key: impl AsRef<str>, value: Self::Value) -> Self {
81 Self(format!("n.{} > {}", key.as_ref(), serialize_cypher(value)))
82 }
83
84 fn lt(key: impl AsRef<str>, value: Self::Value) -> Self {
85 Self(format!("n.{} < {}", key.as_ref(), serialize_cypher(value)))
86 }
87
88 fn and(self, rhs: Self) -> Self {
89 Self(format!("({}) AND ({})", self.0, rhs.0))
90 }
91
92 fn or(self, rhs: Self) -> Self {
93 Self(format!("({}) OR ({})", self.0, rhs.0))
94 }
95}
96
97impl Neo4jSearchFilter {
98 pub fn render(self) -> String {
99 format!("WHERE {}", self.0)
100 }
101
102 pub fn not(self) -> Self {
103 Self(format!("NOT ({})", self.0))
104 }
105
106 pub fn gte(key: &str, value: <Self as SearchFilter>::Value) -> Self {
107 Self(format!("n.{key} >= {}", serialize_cypher(value)))
108 }
109
110 pub fn lte(key: &str, value: <Self as SearchFilter>::Value) -> Self {
111 Self(format!("n.{key} <= {}", serialize_cypher(value)))
112 }
113
114 pub fn member(key: &str, values: Vec<<Self as SearchFilter>::Value>) -> Self {
115 Self(format!(
116 "n.{key} IN {}",
117 serialize_cypher(serde_json::Value::Array(values))
118 ))
119 }
120
121 pub fn contains<S>(key: &str, pattern: S) -> Self
123 where
124 S: AsRef<str>,
125 {
126 Self(format!(
127 "n.{key} CONTAINS {}",
128 serialize_cypher(serde_json::Value::String(pattern.as_ref().into()))
129 ))
130 }
131
132 pub fn starts_with<S>(key: &str, pattern: S) -> Self
134 where
135 S: AsRef<str>,
136 {
137 Self(format!(
138 "n.{key} STARTS WITH {}",
139 serialize_cypher(serde_json::Value::String(pattern.as_ref().into()))
140 ))
141 }
142
143 pub fn ends_with<S>(key: &str, pattern: S) -> Self
145 where
146 S: AsRef<str>,
147 {
148 Self(format!(
149 "n.{key} ENDS WITH {}",
150 serialize_cypher(serde_json::Value::String(pattern.as_ref().into()))
151 ))
152 }
153
154 pub fn matches<S>(key: &str, pattern: S) -> Self
156 where
157 S: AsRef<str>,
158 {
159 Self(format!(
160 "n.{key} =~ {}",
161 serialize_cypher(serde_json::Value::String(pattern.as_ref().into()))
162 ))
163 }
164}
165
166fn serialize_cypher(value: serde_json::Value) -> String {
168 use serde_json::Value::*;
169 match value {
170 Null => "null".into(),
171 Bool(b) => b.to_string(),
172 Number(n) => n.to_string(),
173 String(s) => format!("'{}'", s.replace('\'', "\\'")),
174 Array(arr) => {
175 format!(
176 "[{}]",
177 arr.into_iter()
178 .map(serialize_cypher)
179 .collect::<Vec<std::string::String>>()
180 .join(", ")
181 )
182 }
183 Object(obj) => {
184 format!(
185 "{{{}}}",
186 obj.into_iter()
187 .map(|(k, v)| format!("{k}: {}", serialize_cypher(v)))
188 .collect::<Vec<std::string::String>>()
189 .join(", ")
190 )
191 }
192 }
193}
194
195pub trait ToBoltType {
197 fn to_bolt_type(&self) -> BoltType;
200}
201
202impl<T> ToBoltType for T
203where
204 T: serde::Serialize,
205{
206 fn to_bolt_type(&self) -> BoltType {
207 match serde_json::to_value(self) {
208 Ok(json_value) => match json_value {
209 serde_json::Value::Null => BoltType::Null(BoltNull),
210 serde_json::Value::Bool(b) => BoltType::Boolean(BoltBoolean::new(b)),
211 serde_json::Value::Number(num) => {
212 if let Some(i) = num.as_i64() {
213 BoltType::Integer(BoltInteger::new(i))
214 } else if let Some(f) = num.as_f64() {
215 BoltType::Float(BoltFloat::new(f))
216 } else {
217 println!("Couldn't map to BoltType, will ignore.");
218 BoltType::Null(BoltNull)
219 }
220 }
221 serde_json::Value::String(s) => BoltType::String(BoltString::new(&s)),
222 serde_json::Value::Array(arr) => BoltType::List(
223 arr.iter()
224 .map(ToBoltType::to_bolt_type)
225 .collect::<Vec<BoltType>>()
226 .into(),
227 ),
228 serde_json::Value::Object(obj) => {
229 let mut bolt_map = BoltMap::new();
230 for (k, v) in obj {
231 bolt_map.put(BoltString::new(&k), v.to_bolt_type());
232 }
233 BoltType::Map(bolt_map)
234 }
235 },
236 Err(_) => {
237 println!("Couldn't serialize to JSON, will ignore.");
238 BoltType::Null(BoltNull)
239 }
240 }
241 }
242}
243
244impl Neo4jClient {
245 const GET_INDEX_QUERY: &'static str = "
246 SHOW VECTOR INDEXES
247 YIELD name, labelsOrTypes, properties, options
248 WHERE name=$index_name
249 RETURN name, labelsOrTypes, properties, options
250 ";
251
252 const SHOW_INDEXES_QUERY: &'static str = "SHOW VECTOR INDEXES YIELD name RETURN name";
253
254 pub fn new(graph: Graph) -> Self {
255 Self { graph }
256 }
257
258 pub async fn connect(uri: &str, user: &str, password: &str) -> Result<Self, VectorStoreError> {
259 tracing::info!("Connecting to Neo4j DB at {} ...", uri);
260 let graph = Graph::new(uri, user, password)
261 .await
262 .map_err(VectorStoreError::datastore)?;
263 tracing::info!("Connected to Neo4j");
264 Ok(Self { graph })
265 }
266
267 pub async fn from_config(config: Config) -> Result<Self, VectorStoreError> {
268 let graph = Graph::connect(config)
269 .await
270 .map_err(VectorStoreError::datastore)?;
271 Ok(Self { graph })
272 }
273
274 pub async fn execute_and_collect<T: for<'a> Deserialize<'a>>(
275 graph: &Graph,
276 query: Query,
277 ) -> Result<Vec<T>, VectorStoreError> {
278 graph
279 .execute(query)
280 .await
281 .map_err(VectorStoreError::datastore)?
282 .into_stream_as::<T>()
283 .try_collect::<Vec<T>>()
284 .await
285 .map_err(VectorStoreError::datastore)
286 }
287
288 pub async fn get_index(
295 &self,
296 model: impl Into<rig_core::DynModel<rig_core::operation::Embedding>>,
297 index_name: &str,
298 ) -> Result<Neo4jVectorIndex, VectorStoreError> {
299 let model: rig_core::DynModel<rig_core::operation::Embedding> = model.into();
300 #[derive(Deserialize)]
301 #[serde(rename_all = "camelCase")]
302 struct IndexInfo {
303 name: String,
304 labels_or_types: Vec<String>,
305 properties: Vec<String>,
306 options: IndexOptions,
307 }
308
309 #[derive(Deserialize)]
310 #[serde(rename_all = "camelCase")]
311 struct IndexOptions {
312 #[allow(dead_code)]
313 index_provider: Option<String>,
314 index_config: IndexConfigDetails,
315 }
316
317 #[derive(Deserialize)]
318 struct IndexConfigDetails {
319 #[serde(rename = "vector.dimensions")]
320 vector_dimensions: i64,
321 #[serde(rename = "vector.similarity_function")]
322 vector_similarity_function: String,
323 }
324
325 let index_info = Self::execute_and_collect::<IndexInfo>(
326 &self.graph,
327 neo4rs::query(Self::GET_INDEX_QUERY).param("index_name", index_name),
328 )
329 .await?;
330
331 let index_config = if let Some(index) = index_info.first() {
332 if index.options.index_config.vector_dimensions != model.capabilities().ndims as i64 {
333 tracing::warn!(
334 "The embedding vector dimensions of the existing Neo4j DB index ({}) do not match the provided model dimensions ({}). This may affect search performance.",
335 index.options.index_config.vector_dimensions,
336 model.capabilities().ndims
337 );
338 }
339 let embedding_property = index.properties.first().ok_or_else(|| {
340 VectorStoreError::DatastoreError(
341 "Neo4j index is missing an embedding property".into(),
342 )
343 })?;
344 let mut config = IndexConfig::new(index.name.clone())
345 .embedding_property(embedding_property)
346 .similarity_function(VectorSimilarityFunction::from_str(
347 &index.options.index_config.vector_similarity_function,
348 )?);
349 if let Some(label) = index.labels_or_types.first() {
351 config = config.node_label(label);
352 }
353 config
354 } else {
355 let indexes = Self::execute_and_collect::<String>(
356 &self.graph,
357 neo4rs::query(Self::SHOW_INDEXES_QUERY),
358 )
359 .await?;
360 return Err(VectorStoreError::datastore(std::io::Error::new(
361 std::io::ErrorKind::NotFound,
362 format!(
363 "Index `{index_name}` not found in database. Available indexes: {indexes:?}"
364 ),
365 )));
366 };
367 Ok(Neo4jVectorIndex::new(
368 self.graph.clone(),
369 model,
370 index_config,
371 ))
372 }
373
374 pub async fn create_vector_index(
385 &self,
386 index_config: IndexConfig,
387 node_label: &str,
388 ndims: usize,
389 ) -> Result<(), VectorStoreError> {
390 tracing::info!("Creating vector index {} ...", index_config.index_name);
391
392 let create_vector_index_query = format!(
393 "
394 CREATE VECTOR INDEX $index_name IF NOT EXISTS
395 FOR (m:{})
396 ON m.{}
397 OPTIONS {{
398 indexConfig: {{
399 `vector.dimensions`: $dimensions,
400 `vector.similarity_function`: $similarity_function
401 }}
402 }}",
403 node_label, index_config.embedding_property
404 );
405
406 self.graph
407 .run(
408 neo4rs::query(&create_vector_index_query)
409 .param("index_name", index_config.index_name.clone())
410 .param(
411 "similarity_function",
412 index_config.similarity_function.clone().to_bolt_type(),
413 )
414 .param("dimensions", ndims as i64),
415 )
416 .await
417 .map_err(VectorStoreError::datastore)?;
418
419 let index_exists = self
420 .graph
421 .run(
422 neo4rs::query("CALL db.awaitIndex($index_name, 10000)")
423 .param("index_name", index_config.index_name.clone()),
424 )
425 .await;
426
427 if index_exists.is_err() {
428 tracing::warn!(
429 "Index with name `{}` is not ready or could not be created.",
430 index_config.index_name.clone()
431 );
432 }
433
434 tracing::info!(
435 "Index created successfully with name: {}",
436 index_config.index_name
437 );
438 Ok(())
439 }
440}