Skip to main content

rig_neo4j/
lib.rs

1//! Neo4j vector store for Rig.
2//!
3//! [`Neo4jVectorIndex`] queries a vector index that must already exist, created
4//! externally or through [`Neo4jClient::create_vector_index`]. Neo4j builds new
5//! indexes in the background, so they are not queryable immediately. Self-managed
6//! instances need the GenAI plugin installed; Neo4j Aura enables it by default.
7//! The crate [README](https://github.com/0xPlaygrounds/rig/tree/main/crates/rig-neo4j)
8//! covers setup and further examples.
9//!
10//! ```no_run
11//! use neo4rs::ConfigBuilder;
12//! use rig_core::providers::openai::{self, OpenAI};
13//! use rig_core::vector_store::VectorStoreIndex;
14//! use rig_core::vector_store::request::VectorSearchRequest;
15//! use rig_neo4j::Neo4jClient;
16//! use serde::Deserialize;
17//!
18//! #[derive(Debug, Deserialize)]
19//! struct Movie {
20//!     title: String,
21//!     plot: String,
22//! }
23//!
24//! #[tokio::main]
25//! async fn main() -> Result<(), anyhow::Error> {
26//!     let openai = OpenAI::from_env()?;
27//!     let model = openai.embedding(openai::TEXT_EMBEDDING_ADA_002, None);
28//!
29//!     let client = Neo4jClient::from_config(
30//!         ConfigBuilder::default()
31//!             .uri("neo4j+s://demo.neo4jlabs.com:7687")
32//!             .db("recommendations")
33//!             .user("recommendations")
34//!             .password("recommendations")
35//!             .build()?,
36//!     )
37//!     .await?;
38//!
39//!     // ❗IMPORTANT: reuse the model the stored embeddings were generated with.
40//!     let index = client.get_index(model, "moviePlotsEmbedding").await?;
41//!
42//!     let req = VectorSearchRequest::builder()
43//!         .query("Batman")
44//!         .samples(3)
45//!         .build();
46//!
47//!     let results = index.top_n::<Movie>(req).await?;
48//!     println!("{results:#?}");
49//!
50//!     Ok(())
51//! }
52//! ```
53pub 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/// Cypher predicate over the matched node `n`.
67///
68/// Property keys are spliced into the query verbatim and string values are only
69/// single-quote escaped, so neither should carry untrusted input.
70#[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    /// Matches property values containing `pattern`.
122    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    /// Matches property values starting with `pattern`.
133    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    /// Matches property values ending with `pattern`.
144    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    /// Matches property values against the Cypher regular expression `pattern`.
155    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
166/// Renders a JSON value as a Cypher literal, escaping single quotes in strings.
167fn 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
195/// Conversion into a Bolt parameter value.
196pub trait ToBoltType {
197    /// Converts through JSON, yielding `BoltType::Null` for values that fail to
198    /// serialize or fall outside Bolt's numeric range.
199    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    /// Returns an index handle mirroring the existing vector index `index_name`,
289    /// adopting its embedding property, similarity function, and node label.
290    ///
291    /// `model` must be the model whose embeddings populated the index; a
292    /// dimension mismatch is only warned about. Errors when the index does not
293    /// exist or defines no property.
294    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            // Inserts must target the label the index is attached to.
350            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    /// Creates a vector index over `node_label` if one of that name does not
375    /// already exist, over vectors of `ndims` dimensions.
376    ///
377    /// `ndims` must be the width of the model later handed to
378    /// [`Self::get_index`], read as `capabilities().ndims`; an index built at
379    /// one width and queried with another returns meaningless results.
380    ///
381    /// `node_label` and the configured embedding property are spliced into the
382    /// Cypher statement verbatim. Waiting for the index to come online is
383    /// best effort: a timeout is logged as a warning rather than returned.
384    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}