diff --git a/docs/open-source/features/graph-memory.mdx b/docs/open-source/features/graph-memory.mdx
index f8c24bbc9..c7499d4de 100644
--- a/docs/open-source/features/graph-memory.mdx
+++ b/docs/open-source/features/graph-memory.mdx
@@ -35,7 +35,7 @@ graph LR
Mem0’s extraction LLM identifies entities, relationships, and timestamps from the conversation payload you send to `memory.add`.
-Embeddings land in your configured vector database while nodes and edges flow into a Bolt-compatible graph backend (Neo4j, Memgraph, Neptune, or Kuzu).
+Embeddings land in your configured vector database while nodes and edges flow into a graph backend (Neo4j, Memgraph, Neptune, Kuzu, or Apache AGE).
`memory.search` performs vector similarity (optionally reranked by your configured reranker) and returns the results list. Graph Memory runs in parallel and adds related entities in the `relations` array—it does not reorder the vector hits automatically.
@@ -264,7 +264,7 @@ Monitor graph growth, especially on free tiers, by periodically cleaning dormant
## Decision Points
-- Select the graph store that fits your deployment (managed Aura vs. self-hosted Neo4j vs. AWS Neptune vs. local Kuzu).
+- Select the graph store that fits your deployment (managed Aura vs. self-hosted Neo4j vs. AWS Neptune vs. local Kuzu vs. Apache AGE on PostgreSQL).
- Decide when to enable graph writes per request; routine conversations may stay vector-only to save latency.
- Set a policy for pruning stale relationships so your graph stays fast and affordable.
@@ -381,6 +381,42 @@ config = {
Kuzu will clear its state when using `:memory:` once the process exits. See the [Kuzu documentation](https://kuzudb.com/docs/) for advanced settings.
+
+ [Apache AGE](https://age.apache.org/) adds graph database capabilities to PostgreSQL, letting you run Cypher queries alongside SQL on the same server. Start AGE via Docker, then configure Mem0:
+
+```bash
+docker run --name age-postgres \
+ -e POSTGRES_DB=mem0_db \
+ -e POSTGRES_USER=mem0_user \
+ -e POSTGRES_PASSWORD=mem0_pass \
+ -p 5432:5432 \
+ -d apache/age
+```
+
+```python
+from mem0 import Memory
+
+config = {
+ "graph_store": {
+ "provider": "apache_age",
+ "config": {
+ "host": "localhost",
+ "port": 5432,
+ "database": "mem0_db",
+ "username": "mem0_user",
+ "password": "mem0_pass",
+ "graph_name": "mem0_graph",
+ },
+ },
+}
+
+m = Memory.from_config(config_dict=config)
+```
+
+Apache AGE does not have a built-in vector index, so similarity search is computed client-side. This works well for moderate graph sizes; for very large graphs consider pairing AGE with pgvector for the vector store.
+
+Reference: [Apache AGE documentation](https://age.apache.org/age-manual/master/index.html).
+
diff --git a/mem0/graphs/configs.py b/mem0/graphs/configs.py
index 19bb17c40..b2a7b9ef8 100644
--- a/mem0/graphs/configs.py
+++ b/mem0/graphs/configs.py
@@ -77,12 +77,32 @@ class KuzuConfig(BaseModel):
db: Optional[str] = Field(":memory:", description="Path to a Kuzu database file")
+class ApacheAgeConfig(BaseModel):
+ host: Optional[str] = Field("localhost", description="PostgreSQL server hostname")
+ port: Optional[int] = Field(5432, description="PostgreSQL server port")
+ database: Optional[str] = Field(None, description="PostgreSQL database name")
+ username: Optional[str] = Field(None, description="PostgreSQL username")
+ password: Optional[str] = Field(None, description="PostgreSQL password")
+ graph_name: Optional[str] = Field("mem0_graph", description="Name of the Apache AGE graph")
+
+ @model_validator(mode="before")
+ def check_required_fields(cls, values):
+ database, username, password = (
+ values.get("database"),
+ values.get("username"),
+ values.get("password"),
+ )
+ if not database or not username or not password:
+ raise ValueError("Please provide 'database', 'username' and 'password'.")
+ return values
+
+
class GraphStoreConfig(BaseModel):
provider: str = Field(
- description="Provider of the data store (e.g., 'neo4j', 'memgraph', 'neptune', 'kuzu')",
+ description="Provider of the data store (e.g., 'neo4j', 'memgraph', 'neptune', 'kuzu', 'apache_age')",
default="neo4j",
)
- config: Union[Neo4jConfig, MemgraphConfig, NeptuneConfig, KuzuConfig] = Field(
+ config: Union[Neo4jConfig, MemgraphConfig, NeptuneConfig, KuzuConfig, ApacheAgeConfig] = Field(
description="Configuration for the specific data store", default=None
)
llm: Optional[LlmConfig] = Field(description="LLM configuration for querying the graph store", default=None)
@@ -110,5 +130,7 @@ class GraphStoreConfig(BaseModel):
return NeptuneConfig(**v.model_dump())
elif provider == "kuzu":
return KuzuConfig(**v.model_dump())
+ elif provider == "apache_age":
+ return ApacheAgeConfig(**v.model_dump())
else:
raise ValueError(f"Unsupported graph store provider: {provider}")
diff --git a/mem0/memory/apache_age_memory.py b/mem0/memory/apache_age_memory.py
new file mode 100644
index 000000000..3c6117dcc
--- /dev/null
+++ b/mem0/memory/apache_age_memory.py
@@ -0,0 +1,570 @@
+import json
+import logging
+import time
+
+from mem0.memory.utils import format_entities, sanitize_relationship_for_cypher
+
+try:
+ import age
+except ImportError:
+ raise ImportError("apache-age-python is not installed. Please install it using pip install apache-age-python")
+
+try:
+ from rank_bm25 import BM25Okapi
+except ImportError:
+ raise ImportError("rank_bm25 is not installed. Please install it using pip install rank-bm25")
+
+from mem0.graphs.tools import (
+ DELETE_MEMORY_STRUCT_TOOL_GRAPH,
+ DELETE_MEMORY_TOOL_GRAPH,
+ EXTRACT_ENTITIES_STRUCT_TOOL,
+ EXTRACT_ENTITIES_TOOL,
+ RELATIONS_STRUCT_TOOL,
+ RELATIONS_TOOL,
+)
+from mem0.graphs.utils import EXTRACT_RELATIONS_PROMPT, get_delete_messages
+from mem0.utils.factory import EmbedderFactory, LlmFactory
+
+logger = logging.getLogger(__name__)
+
+
+def _cosine_similarity(vec1, vec2):
+ """Compute cosine similarity between two vectors without numpy."""
+ dot = sum(a * b for a, b in zip(vec1, vec2))
+ norm1 = sum(a * a for a in vec1) ** 0.5
+ norm2 = sum(b * b for b in vec2) ** 0.5
+ if norm1 == 0 or norm2 == 0:
+ return 0.0
+ return dot / (norm1 * norm2)
+
+
+class MemoryGraph:
+ def __init__(self, config):
+ self.config = config
+
+ graph_cfg = self.config.graph_store.config
+ self.graph_name = graph_cfg.graph_name
+
+ # Connect using the Apache AGE Python driver (psycopg2-based)
+ self.ag = age.connect(
+ graph=self.graph_name,
+ host=graph_cfg.host,
+ port=graph_cfg.port,
+ dbname=graph_cfg.database,
+ user=graph_cfg.username,
+ password=graph_cfg.password,
+ )
+
+ self.embedding_model = EmbedderFactory.create(
+ self.config.embedder.provider, self.config.embedder.config, self.config.vector_store.config
+ )
+
+ # Default to openai if no specific provider is configured
+ self.llm_provider = "openai"
+ if self.config.llm and self.config.llm.provider:
+ self.llm_provider = self.config.llm.provider
+ if self.config.graph_store and self.config.graph_store.llm and self.config.graph_store.llm.provider:
+ self.llm_provider = self.config.graph_store.llm.provider
+
+ # Get LLM config with proper null checks
+ llm_config = None
+ if self.config.graph_store and self.config.graph_store.llm and hasattr(self.config.graph_store.llm, "config"):
+ llm_config = self.config.graph_store.llm.config
+ elif hasattr(self.config.llm, "config"):
+ llm_config = self.config.llm.config
+ self.llm = LlmFactory.create(self.llm_provider, llm_config)
+ self.user_id = None
+ self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, "threshold") else 0.7
+
+ # -- helpers ---------------------------------------------------------------
+
+ def _exec_cypher(self, cypher_stmt, cols=None, params=None):
+ """Execute a Cypher query via the AGE driver and return results.
+
+ Uses ``ag.execCypher`` which delegates to psycopg2's safe parameter
+ substitution (``%s`` placeholders). The *cols* argument specifies the
+ column names in the ``AS (…)`` clause — when ``None`` the driver
+ defaults to a single ``v agtype`` column.
+
+ When *cols* are provided, returns a list of dicts keyed by column name.
+ When *cols* is ``None``, returns vertex/edge property dicts or raw values.
+ """
+ cursor = self.ag.execCypher(cypher_stmt, cols=cols, params=params)
+ rows = cursor.fetchall()
+ if not rows:
+ return []
+
+ col_names = [desc[0] for desc in cursor.description] if cursor.description else None
+ results = []
+ for row in rows:
+ if col_names and len(col_names) > 1:
+ record = {}
+ for i, col_name in enumerate(col_names):
+ val = row[i]
+ if hasattr(val, "properties"):
+ record[col_name] = val.properties
+ else:
+ record[col_name] = val
+ results.append(record)
+ else:
+ val = row[0] if len(row) == 1 else row
+ if hasattr(val, "properties"):
+ results.append(val.properties)
+ else:
+ results.append(val)
+ return results
+
+ def _find_similar_node(self, embedding, filters, threshold=0.9):
+ """Find the most similar existing node by cosine similarity.
+
+ Apache AGE does not have a built-in vector index, so we fetch all
+ node embeddings matching the filters and compute cosine similarity
+ on the client side. This is adequate for moderate graph sizes; for
+ very large graphs consider pairing AGE with pgvector.
+ """
+ nodes = self._exec_cypher(
+ "MATCH (n {user_id: %s}) WHERE n.embedding IS NOT NULL RETURN n",
+ params=(filters["user_id"],),
+ )
+
+ best_match = None
+ best_sim = 0.0
+ for node in nodes:
+ props = node if isinstance(node, dict) else {}
+ stored_emb = props.get("embedding")
+ if not stored_emb:
+ continue
+ if isinstance(stored_emb, str):
+ stored_emb = json.loads(stored_emb)
+
+ if filters.get("agent_id") and props.get("agent_id") != filters["agent_id"]:
+ continue
+ if filters.get("run_id") and props.get("run_id") != filters["run_id"]:
+ continue
+
+ sim = _cosine_similarity(embedding, stored_emb)
+ if sim >= threshold and sim > best_sim:
+ best_sim = sim
+ best_match = props
+
+ return best_match
+
+ def _merge_node(self, user_id, name, embedding, agent_id=None, run_id=None):
+ """Create a node if it doesn't exist, or update mentions if it does.
+
+ Apache AGE does not support ``ON CREATE SET`` / ``ON MATCH SET``, so
+ we use ``MERGE … SET`` which always applies the SET clause. Embeddings
+ are always refreshed on merge to keep them current.
+ """
+ self._exec_cypher(
+ "MERGE (n {user_id: %s, name: %s}) SET n.embedding = %s, "
+ "n.mentions = coalesce(n.mentions, 0) + 1, "
+ "n.created = coalesce(n.created, %s)",
+ params=(user_id, name, json.dumps(embedding), int(time.time() * 1000)),
+ )
+ if agent_id:
+ self._exec_cypher(
+ "MATCH (n {user_id: %s, name: %s}) SET n.agent_id = %s",
+ params=(user_id, name, agent_id),
+ )
+ if run_id:
+ self._exec_cypher(
+ "MATCH (n {user_id: %s, name: %s}) SET n.run_id = %s",
+ params=(user_id, name, run_id),
+ )
+
+ # -- public API ------------------------------------------------------------
+
+ def add(self, data, filters):
+ """
+ Adds data to the graph.
+
+ Args:
+ data (str): The data to add to the graph.
+ filters (dict): A dictionary containing filters to be applied during the addition.
+ """
+ entity_type_map = self._retrieve_nodes_from_data(data, filters)
+ to_be_added = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
+ search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
+ to_be_deleted = self._get_delete_entities_from_search_output(search_output, data, filters)
+
+ deleted_entities = self._delete_entities(to_be_deleted, filters)
+ added_entities = self._add_entities(to_be_added, filters, entity_type_map)
+
+ return {"deleted_entities": deleted_entities, "added_entities": added_entities}
+
+ def search(self, query, filters, limit=100):
+ """
+ Search for memories and related graph data.
+
+ Args:
+ query (str): Query to search for.
+ filters (dict): A dictionary containing filters to be applied during the search.
+ limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
+
+ Returns:
+ dict: A dictionary containing:
+ - "contexts": List of search results from the base data store.
+ - "entities": List of related graph data based on the query.
+ """
+ entity_type_map = self._retrieve_nodes_from_data(query, filters)
+ search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
+
+ if not search_output:
+ return []
+
+ search_outputs_sequence = [
+ [item["source"], item["relationship"], item["destination"]] for item in search_output
+ ]
+ bm25 = BM25Okapi(search_outputs_sequence)
+
+ tokenized_query = query.split(" ")
+ reranked_results = bm25.get_top_n(tokenized_query, search_outputs_sequence, n=5)
+
+ search_results = []
+ for item in reranked_results:
+ search_results.append({"source": item[0], "relationship": item[1], "destination": item[2]})
+
+ logger.info(f"Returned {len(search_results)} search results")
+ return search_results
+
+ def delete_all(self, filters):
+ """Delete all nodes and relationships for a user or specific agent."""
+ # AGE doesn't support compound property maps in MATCH the same way
+ # Neo4j does, so we use a WHERE clause to handle optional filters.
+ where_parts = ["n.user_id = %s"]
+ params = [filters["user_id"]]
+ if filters.get("agent_id"):
+ where_parts.append("n.agent_id = %s")
+ params.append(filters["agent_id"])
+ if filters.get("run_id"):
+ where_parts.append("n.run_id = %s")
+ params.append(filters["run_id"])
+ where_clause = " AND ".join(where_parts)
+
+ self._exec_cypher(
+ f"MATCH (n) WHERE {where_clause} DETACH DELETE n",
+ params=tuple(params),
+ )
+ self.ag.commit()
+
+ def get_all(self, filters, limit=100):
+ """
+ Retrieves all nodes and relationships from the graph database based on optional filtering criteria.
+
+ Args:
+ filters (dict): A dictionary containing filters to be applied during the retrieval.
+ limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
+ Returns:
+ list: A list of dictionaries, each containing:
+ - 'source': The source node name.
+ - 'relationship': The relationship type.
+ - 'target': The target node name.
+ """
+ where_parts = ["n.user_id = %s", "m.user_id = %s"]
+ params = [filters["user_id"], filters["user_id"]]
+ if filters.get("agent_id"):
+ where_parts.extend(["n.agent_id = %s", "m.agent_id = %s"])
+ params.extend([filters["agent_id"], filters["agent_id"]])
+ if filters.get("run_id"):
+ where_parts.extend(["n.run_id = %s", "m.run_id = %s"])
+ params.extend([filters["run_id"], filters["run_id"]])
+ where_clause = " AND ".join(where_parts)
+
+ results = self._exec_cypher(
+ f"MATCH (n)-[r]->(m) WHERE {where_clause} "
+ f"RETURN n.name, type(r), m.name",
+ cols=["source", "relationship", "target"],
+ params=tuple(params),
+ )
+
+ final_results = []
+ for result in results[:limit]:
+ final_results.append(
+ {
+ "source": result["source"],
+ "relationship": result["relationship"],
+ "target": result["target"],
+ }
+ )
+
+ logger.info(f"Retrieved {len(final_results)} relationships")
+ return final_results
+
+ # -- LLM-driven extraction -------------------------------------------------
+
+ def _retrieve_nodes_from_data(self, data, filters):
+ """Extracts all the entities mentioned in the query."""
+ _tools = [EXTRACT_ENTITIES_TOOL]
+ if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
+ _tools = [EXTRACT_ENTITIES_STRUCT_TOOL]
+ search_results = self.llm.generate_response(
+ messages=[
+ {
+ "role": "system",
+ "content": f"You are a smart assistant who understands entities and their types in a given text. If user message contains self reference such as 'I', 'me', 'my' etc. then use {filters['user_id']} as the source entity. Extract all the entities from the text. ***DO NOT*** answer the question itself if the given text is a question.",
+ },
+ {"role": "user", "content": data},
+ ],
+ tools=_tools,
+ )
+
+ entity_type_map = {}
+
+ try:
+ for tool_call in search_results["tool_calls"]:
+ if tool_call["name"] != "extract_entities":
+ continue
+ for item in tool_call.get("arguments", {}).get("entities", []):
+ if "entity" in item and "entity_type" in item:
+ entity_type_map[item["entity"]] = item["entity_type"]
+ except Exception as e:
+ logger.exception(
+ f"Error in search tool: {e}, llm_provider={self.llm_provider}, search_results={search_results}"
+ )
+
+ entity_type_map = {k.lower().replace(" ", "_"): v.lower().replace(" ", "_") for k, v in entity_type_map.items()}
+ logger.debug(f"Entity type map: {entity_type_map}\n search_results={search_results}")
+ return entity_type_map
+
+ def _establish_nodes_relations_from_data(self, data, filters, entity_type_map):
+ """Establish relations among the extracted nodes."""
+
+ user_identity = f"user_id: {filters['user_id']}"
+ if filters.get("agent_id"):
+ user_identity += f", agent_id: {filters['agent_id']}"
+ if filters.get("run_id"):
+ user_identity += f", run_id: {filters['run_id']}"
+
+ if self.config.graph_store.custom_prompt:
+ system_content = EXTRACT_RELATIONS_PROMPT.replace("USER_ID", user_identity)
+ system_content = system_content.replace("CUSTOM_PROMPT", f"4. {self.config.graph_store.custom_prompt}")
+ messages = [
+ {"role": "system", "content": system_content},
+ {"role": "user", "content": data},
+ ]
+ else:
+ system_content = EXTRACT_RELATIONS_PROMPT.replace("USER_ID", user_identity)
+ messages = [
+ {"role": "system", "content": system_content},
+ {"role": "user", "content": f"List of entities: {list(entity_type_map.keys())}. \n\nText: {data}"},
+ ]
+
+ _tools = [RELATIONS_TOOL]
+ if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
+ _tools = [RELATIONS_STRUCT_TOOL]
+
+ extracted_entities = self.llm.generate_response(
+ messages=messages,
+ tools=_tools,
+ )
+
+ entities = []
+ if extracted_entities and extracted_entities.get("tool_calls"):
+ entities = extracted_entities["tool_calls"][0].get("arguments", {}).get("entities", [])
+
+ entities = self._remove_spaces_from_entities(entities)
+ logger.debug(f"Extracted entities: {entities}")
+ return entities
+
+ # -- graph DB operations ---------------------------------------------------
+
+ def _search_graph_db(self, node_list, filters, limit=100):
+ """Search similar nodes and their respective incoming and outgoing relations."""
+ result_relations = []
+
+ for node in node_list:
+ n_embedding = self.embedding_model.embed(node)
+
+ # Fetch all nodes with embeddings for this user
+ nodes = self._exec_cypher(
+ "MATCH (n {user_id: %s}) WHERE n.embedding IS NOT NULL RETURN n",
+ params=(filters["user_id"],),
+ )
+
+ # Compute cosine similarity client-side and collect similar nodes
+ similar_nodes = []
+ for node_result in nodes:
+ props = node_result if isinstance(node_result, dict) else {}
+ stored_emb = props.get("embedding")
+ if not stored_emb:
+ continue
+ if isinstance(stored_emb, str):
+ stored_emb = json.loads(stored_emb)
+
+ if filters.get("agent_id") and props.get("agent_id") != filters["agent_id"]:
+ continue
+ if filters.get("run_id") and props.get("run_id") != filters["run_id"]:
+ continue
+
+ sim = _cosine_similarity(n_embedding, stored_emb)
+ if sim >= self.threshold:
+ similar_nodes.append({"name": props.get("name"), "similarity": sim})
+
+ similar_nodes.sort(key=lambda x: x["similarity"], reverse=True)
+
+ # Build WHERE clause for relationship target filtering
+ rel_where_parts = ["m.user_id = %s"]
+ rel_params_suffix = [filters["user_id"]]
+ if filters.get("agent_id"):
+ rel_where_parts.append("m.agent_id = %s")
+ rel_params_suffix.append(filters["agent_id"])
+ if filters.get("run_id"):
+ rel_where_parts.append("m.run_id = %s")
+ rel_params_suffix.append(filters["run_id"])
+ rel_where = " AND ".join(rel_where_parts)
+
+ # For each similar node, fetch its relationships
+ for sn in similar_nodes[:limit]:
+ node_name = sn["name"]
+ similarity = sn["similarity"]
+
+ out_params = (filters["user_id"], node_name) + tuple(rel_params_suffix)
+ out_results = self._exec_cypher(
+ f"MATCH (n {{user_id: %s, name: %s}})-[r]->(m) "
+ f"WHERE {rel_where} "
+ f"RETURN n.name, type(r), m.name",
+ cols=["source", "relationship", "destination"],
+ params=out_params,
+ )
+
+ in_params = (filters["user_id"], node_name) + tuple(rel_params_suffix)
+ in_results = self._exec_cypher(
+ f"MATCH (n {{user_id: %s, name: %s}})<-[r]-(m) "
+ f"WHERE {rel_where} "
+ f"RETURN m.name, type(r), n.name",
+ cols=["source", "relationship", "destination"],
+ params=in_params,
+ )
+
+ for rel in out_results + in_results:
+ rel["similarity"] = similarity
+ result_relations.append(rel)
+
+ return result_relations
+
+ def _get_delete_entities_from_search_output(self, search_output, data, filters):
+ """Get the entities to be deleted from the search output."""
+ search_output_string = format_entities(search_output)
+
+ user_identity = f"user_id: {filters['user_id']}"
+ if filters.get("agent_id"):
+ user_identity += f", agent_id: {filters['agent_id']}"
+ if filters.get("run_id"):
+ user_identity += f", run_id: {filters['run_id']}"
+
+ system_prompt, user_prompt = get_delete_messages(search_output_string, data, user_identity)
+
+ _tools = [DELETE_MEMORY_TOOL_GRAPH]
+ if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
+ _tools = [DELETE_MEMORY_STRUCT_TOOL_GRAPH]
+
+ memory_updates = self.llm.generate_response(
+ messages=[
+ {"role": "system", "content": system_prompt},
+ {"role": "user", "content": user_prompt},
+ ],
+ tools=_tools,
+ )
+
+ to_be_deleted = []
+ for item in memory_updates.get("tool_calls", []):
+ if item.get("name") == "delete_graph_memory":
+ to_be_deleted.append(item.get("arguments"))
+ to_be_deleted = self._remove_spaces_from_entities(to_be_deleted)
+ logger.debug(f"Deleted relationships: {to_be_deleted}")
+ return to_be_deleted
+
+ def _delete_entities(self, to_be_deleted, filters):
+ """Delete the entities from the graph."""
+ user_id = filters["user_id"]
+ agent_id = filters.get("agent_id")
+ run_id = filters.get("run_id")
+ results = []
+
+ for item in to_be_deleted:
+ source = item["source"]
+ destination = item["destination"]
+ relationship = item["relationship"]
+
+ where_parts = [
+ "n.user_id = %s", "n.name = %s",
+ "m.user_id = %s", "m.name = %s",
+ ]
+ params = [user_id, source, user_id, destination]
+ if agent_id:
+ where_parts.extend(["n.agent_id = %s", "m.agent_id = %s"])
+ params.extend([agent_id, agent_id])
+ if run_id:
+ where_parts.extend(["n.run_id = %s", "m.run_id = %s"])
+ params.extend([run_id, run_id])
+ where_clause = " AND ".join(where_parts)
+
+ result = self._exec_cypher(
+ f"MATCH (n)-[r:{relationship}]->(m) "
+ f"WHERE {where_clause} "
+ f"DELETE r "
+ f"RETURN n.name, type(r), m.name",
+ cols=["source", "relationship", "target"],
+ params=tuple(params),
+ )
+ results.append(result)
+ self.ag.commit()
+
+ return results
+
+ def _add_entities(self, to_be_added, filters, entity_type_map):
+ """Add new entities to the graph. Merge nodes if they already exist.
+
+ Apache AGE does not support ``ON CREATE SET`` / ``ON MATCH SET``, so we
+ use ``MERGE … SET`` which always applies. The ``coalesce`` pattern
+ ensures ``created`` is only set on the first merge.
+ """
+ user_id = filters["user_id"]
+ agent_id = filters.get("agent_id")
+ run_id = filters.get("run_id")
+ results = []
+
+ for item in to_be_added:
+ source = item["source"]
+ destination = item["destination"]
+ relationship = item["relationship"]
+
+ source_embedding = self.embedding_model.embed(source)
+ dest_embedding = self.embedding_model.embed(destination)
+
+ source_match = self._find_similar_node(source_embedding, filters, threshold=self.threshold)
+ dest_match = self._find_similar_node(dest_embedding, filters, threshold=self.threshold)
+
+ effective_source = source_match["name"] if source_match else source
+ effective_dest = dest_match["name"] if dest_match else destination
+
+ # Merge source and destination nodes
+ self._merge_node(user_id, effective_source, source_embedding, agent_id, run_id)
+ self._merge_node(user_id, effective_dest, dest_embedding, agent_id, run_id)
+
+ # Merge relationship
+ result = self._exec_cypher(
+ f"MATCH (s {{user_id: %s, name: %s}}), (d {{user_id: %s, name: %s}}) "
+ f"MERGE (s)-[r:{relationship}]->(d) "
+ f"RETURN s.name, type(r), d.name",
+ cols=["source", "relationship", "target"],
+ params=(user_id, effective_source, user_id, effective_dest),
+ )
+ self.ag.commit()
+ results.append(result)
+
+ return results
+
+ def _remove_spaces_from_entities(self, entity_list):
+ for item in entity_list:
+ item["source"] = item["source"].lower().replace(" ", "_")
+ item["relationship"] = sanitize_relationship_for_cypher(item["relationship"].lower().replace(" ", "_"))
+ item["destination"] = item["destination"].lower().replace(" ", "_")
+ return entity_list
+
+ def reset(self):
+ """Reset the graph by clearing all nodes and relationships."""
+ logger.warning("Clearing graph...")
+ self._exec_cypher("MATCH (n) DETACH DELETE n")
+ self.ag.commit()
diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py
index ab3fc77a3..3623129d0 100644
--- a/mem0/utils/factory.py
+++ b/mem0/utils/factory.py
@@ -216,6 +216,7 @@ class GraphStoreFactory:
"neptune": "mem0.graphs.neptune.neptunegraph.MemoryGraph",
"neptunedb": "mem0.graphs.neptune.neptunedb.MemoryGraph",
"kuzu": "mem0.memory.kuzu_memory.MemoryGraph",
+ "apache_age": "mem0.memory.apache_age_memory.MemoryGraph",
"default": "mem0.memory.graph_memory.MemoryGraph",
}
diff --git a/pyproject.toml b/pyproject.toml
index 036819381..38a07433c 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -31,6 +31,7 @@ graph = [
"neo4j>=5.23.1",
"rank-bm25>=0.2.2",
"kuzu>=0.11.0",
+ "apache-age-python>=0.0.6",
]
vector_stores = [
"vecs>=0.4.0",
diff --git a/tests/memory/test_apache_age_e2e.py b/tests/memory/test_apache_age_e2e.py
new file mode 100644
index 000000000..99dc3c5a1
--- /dev/null
+++ b/tests/memory/test_apache_age_e2e.py
@@ -0,0 +1,1105 @@
+"""End-to-end integration tests for Apache AGE graph memory.
+
+These tests run against a real Apache AGE instance (via Docker) and exercise
+every layer of the MemoryGraph class: connection, node MERGE, relationship
+MERGE, embedding storage/retrieval, similarity search, deletion, and the
+full add/search/get_all/delete_all/reset public API.
+
+Requirements:
+ docker run --name age-test \
+ -e POSTGRES_DB=mem0_test -e POSTGRES_USER=mem0_user \
+ -e POSTGRES_PASSWORD=mem0_pass -p 15432:5432 -d apache/age
+
+Run:
+ pytest tests/memory/test_apache_age_e2e.py -v -s
+"""
+
+import json
+import os
+import pytest
+from unittest.mock import MagicMock, patch
+
+import age
+
+from mem0.memory.apache_age_memory import MemoryGraph, _cosine_similarity
+
+# -- E2E test configuration ---------------------------------------------------
+
+AGE_HOST = os.environ.get("AGE_HOST", "localhost")
+AGE_PORT = int(os.environ.get("AGE_PORT", "15432"))
+AGE_DB = os.environ.get("AGE_DB", "mem0_test")
+AGE_USER = os.environ.get("AGE_USER", "mem0_user")
+AGE_PASS = os.environ.get("AGE_PASS", "mem0_pass")
+GRAPH_NAME = "e2e_test_graph"
+
+
+def _age_available():
+ """Check if the AGE database is reachable."""
+ try:
+ ag = age.connect(
+ graph=GRAPH_NAME,
+ host=AGE_HOST, port=AGE_PORT,
+ dbname=AGE_DB, user=AGE_USER, password=AGE_PASS,
+ )
+ ag.close()
+ return True
+ except Exception:
+ return False
+
+
+skip_no_age = pytest.mark.skipif(
+ not _age_available(),
+ reason="Apache AGE not available (start Docker container first)",
+)
+
+
+# -- Helpers -------------------------------------------------------------------
+
+def _make_e2e_instance(graph_name=GRAPH_NAME):
+ """Create a MemoryGraph instance wired to the real AGE database,
+ but with LLM and embedding model mocked out."""
+ with patch.object(MemoryGraph, "__init__", return_value=None):
+ mg = MemoryGraph.__new__(MemoryGraph)
+
+ mg.ag = age.connect(
+ graph=graph_name,
+ host=AGE_HOST, port=AGE_PORT,
+ dbname=AGE_DB, user=AGE_USER, password=AGE_PASS,
+ )
+ mg.graph_name = graph_name
+ mg.threshold = 0.7
+
+ # Mock LLM — not needed for DB-layer tests
+ mg.llm_provider = "openai"
+ mg.llm = MagicMock()
+ mg.config = MagicMock()
+ mg.config.graph_store.custom_prompt = None
+
+ # Deterministic embedding model: return a fixed vector derived from the
+ # entity name so similarity searches are predictable. Uses a longer
+ # vector (16-dim) with multiple hash seeds to minimize collisions.
+ def _fake_embed(text):
+ """Map text to a deterministic 16-dim vector for testing."""
+ import hashlib
+ digest = hashlib.sha256(text.encode()).digest()
+ return [b / 255.0 for b in digest[:16]]
+
+ mg.embedding_model = MagicMock()
+ mg.embedding_model.embed = _fake_embed
+ mg.user_id = None
+
+ return mg
+
+
+def _cleanup(mg):
+ """Remove all nodes and close the connection."""
+ try:
+ mg._exec_cypher("MATCH (n) DETACH DELETE n")
+ mg.ag.commit()
+ except Exception:
+ pass
+ try:
+ mg.ag.close()
+ except Exception:
+ pass
+
+
+# ==============================================================================
+# Test: Low-level _exec_cypher
+# ==============================================================================
+
+@skip_no_age
+class TestExecCypher:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_create_and_return_vertex(self):
+ results = self.mg._exec_cypher(
+ "CREATE (n {name: %s, val: %s}) RETURN n",
+ params=("test_node", 42),
+ )
+ self.mg.ag.commit()
+ assert len(results) == 1
+ props = results[0]
+ assert props["name"] == "test_node"
+ assert props["val"] == 42
+
+ def test_return_scalars_with_cols(self):
+ self.mg._exec_cypher(
+ "CREATE (a {name: %s, user_id: %s})", params=("x", "u1")
+ )
+ self.mg._exec_cypher(
+ "CREATE (b {name: %s, user_id: %s})", params=("y", "u1")
+ )
+ self.mg._exec_cypher(
+ "MATCH (a {name: %s}), (b {name: %s}) CREATE (a)-[:LINK]->(b)",
+ params=("x", "y"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg._exec_cypher(
+ "MATCH (n)-[r]->(m) RETURN n.name, type(r), m.name",
+ cols=["source", "rel", "target"],
+ )
+ assert len(results) == 1
+ assert results[0] == {"source": "x", "rel": "LINK", "target": "y"}
+
+ def test_empty_result(self):
+ results = self.mg._exec_cypher(
+ "MATCH (n {name: %s}) RETURN n", params=("nonexistent",)
+ )
+ assert results == []
+
+
+# ==============================================================================
+# Test: _merge_node
+# ==============================================================================
+
+@skip_no_age
+class TestMergeNode:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_creates_node_on_first_merge(self):
+ self.mg._merge_node("u1", "alice", [0.1, 0.2, 0.3])
+ self.mg.ag.commit()
+
+ results = self.mg._exec_cypher(
+ "MATCH (n {user_id: %s, name: %s}) RETURN n",
+ params=("u1", "alice"),
+ )
+ assert len(results) == 1
+ props = results[0]
+ assert props["name"] == "alice"
+ assert props["mentions"] == 1
+ assert props["created"] is not None
+ assert json.loads(props["embedding"]) == [0.1, 0.2, 0.3]
+
+ def test_merge_is_idempotent_increments_mentions(self):
+ self.mg._merge_node("u1", "bob", [0.4, 0.5])
+ self.mg.ag.commit()
+ self.mg._merge_node("u1", "bob", [0.4, 0.5])
+ self.mg.ag.commit()
+
+ results = self.mg._exec_cypher(
+ "MATCH (n {user_id: %s, name: %s}) RETURN n",
+ params=("u1", "bob"),
+ )
+ assert len(results) == 1
+ assert results[0]["mentions"] == 2
+
+ def test_merge_with_agent_id(self):
+ self.mg._merge_node("u1", "carol", [0.6], agent_id="agent_1")
+ self.mg.ag.commit()
+
+ results = self.mg._exec_cypher(
+ "MATCH (n {user_id: %s, name: %s}) RETURN n",
+ params=("u1", "carol"),
+ )
+ assert results[0]["agent_id"] == "agent_1"
+
+
+# ==============================================================================
+# Test: Relationship creation and retrieval
+# ==============================================================================
+
+@skip_no_age
+class TestRelationships:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_create_and_query_relationship(self):
+ self.mg._merge_node("u1", "alice", [0.0]*16)
+ self.mg._merge_node("u1", "bob", [0.0]*16)
+ self.mg.ag.commit()
+
+ self.mg._exec_cypher(
+ "MATCH (s {user_id: %s, name: %s}), (d {user_id: %s, name: %s}) "
+ "MERGE (s)-[r:KNOWS]->(d)",
+ params=("u1", "alice", "u1", "bob"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg._exec_cypher(
+ "MATCH (n {user_id: %s})-[r]->(m) RETURN n.name, type(r), m.name",
+ cols=["source", "rel", "target"],
+ params=("u1",),
+ )
+ assert len(results) == 1
+ assert results[0]["source"] == "alice"
+ assert results[0]["rel"] == "KNOWS"
+ assert results[0]["target"] == "bob"
+
+ def test_multiple_relationships(self):
+ for name in ["alice", "bob", "carol"]:
+ self.mg._merge_node("u1", name, [0.0])
+ self.mg.ag.commit()
+
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:LIKES]->(d)",
+ params=("alice", "carol"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg._exec_cypher(
+ "MATCH (n {name: %s})-[r]->(m) RETURN n.name, type(r), m.name",
+ cols=["source", "rel", "target"],
+ params=("alice",),
+ )
+ assert len(results) == 2
+ rels = {r["rel"] for r in results}
+ assert rels == {"KNOWS", "LIKES"}
+
+
+# ==============================================================================
+# Test: Embedding storage + similarity search
+# ==============================================================================
+
+@skip_no_age
+class TestEmbeddingSimilarity:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_embedding_roundtrip(self):
+ emb = [0.1, 0.2, 0.3, 0.4] + [0.0]*12
+ self.mg._merge_node("u1", "node_a", emb)
+ self.mg.ag.commit()
+
+ results = self.mg._exec_cypher(
+ "MATCH (n {name: %s}) RETURN n", params=("node_a",)
+ )
+ stored = json.loads(results[0]["embedding"])
+ assert stored == emb
+
+ def test_find_similar_node_exact_match(self):
+ emb = [1.0, 0.0, 0.0, 0.0] + [0.0]*12
+ self.mg._merge_node("u1", "target", emb)
+ self.mg.ag.commit()
+
+ match = self.mg._find_similar_node(emb, {"user_id": "u1"}, threshold=0.99)
+ assert match is not None
+ assert match["name"] == "target"
+
+ def test_find_similar_node_no_match_below_threshold(self):
+ self.mg._merge_node("u1", "far_away", [1.0, 0.0, 0.0, 0.0] + [0.0]*12)
+ self.mg.ag.commit()
+
+ orthogonal = [0.0, 1.0, 0.0, 0.0] + [0.0]*12
+ match = self.mg._find_similar_node(orthogonal, {"user_id": "u1"}, threshold=0.5)
+ assert match is None
+
+ def test_find_similar_node_picks_closest(self):
+ # Use vectors where "close" is clearly more similar to the query
+ self.mg._merge_node("u1", "close", [0.9, 0.1, 0.0, 0.0] + [0.0]*12)
+ self.mg._merge_node("u1", "far", [0.0, 0.0, 1.0, 0.0] + [0.0]*12)
+ self.mg.ag.commit()
+
+ query = [1.0, 0.0, 0.0, 0.0] + [0.0]*12
+ match = self.mg._find_similar_node(query, {"user_id": "u1"}, threshold=0.5)
+ assert match is not None
+ assert match["name"] == "close"
+
+ def test_find_similar_node_respects_user_id(self):
+ vec = [1.0, 0.0, 0.0, 0.0] + [0.0]*12
+ self.mg._merge_node("u1", "mine", vec)
+ self.mg._merge_node("u2", "theirs", vec)
+ self.mg.ag.commit()
+
+ match = self.mg._find_similar_node(
+ vec, {"user_id": "u2"}, threshold=0.9
+ )
+ assert match is not None
+ assert match["name"] == "theirs"
+
+
+# ==============================================================================
+# Test: Public API — get_all, delete_all, reset
+# ==============================================================================
+
+@skip_no_age
+class TestPublicAPICRUD:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_get_all_returns_relationships(self):
+ self.mg._merge_node("u1", "alice", [0.0])
+ self.mg._merge_node("u1", "bob", [0.0])
+ self.mg.ag.commit()
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:FRIEND]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg.get_all({"user_id": "u1"})
+ assert len(results) == 1
+ assert results[0]["source"] == "alice"
+ assert results[0]["relationship"] == "FRIEND"
+ assert results[0]["target"] == "bob"
+
+ def test_get_all_empty_for_different_user(self):
+ self.mg._merge_node("u1", "alice", [0.0])
+ self.mg._merge_node("u1", "bob", [0.0])
+ self.mg.ag.commit()
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:FRIEND]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg.get_all({"user_id": "u999"})
+ assert results == []
+
+ def test_get_all_respects_limit(self):
+ for i in range(5):
+ self.mg._merge_node("u1", f"src_{i}", [0.0])
+ self.mg._merge_node("u1", f"dst_{i}", [0.0])
+ self.mg.ag.commit()
+ for i in range(5):
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:REL]->(d)",
+ params=(f"src_{i}", f"dst_{i}"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg.get_all({"user_id": "u1"}, limit=3)
+ assert len(results) == 3
+
+ def test_delete_all_removes_user_data(self):
+ self.mg._merge_node("u1", "alice", [0.0])
+ self.mg._merge_node("u1", "bob", [0.0])
+ self.mg._merge_node("u2", "carol", [0.0])
+ self.mg.ag.commit()
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg.ag.commit()
+
+ self.mg.delete_all({"user_id": "u1"})
+
+ # u1's data should be gone
+ results = self.mg._exec_cypher(
+ "MATCH (n {user_id: %s}) RETURN n", params=("u1",)
+ )
+ assert results == []
+
+ # u2's data should still exist
+ results = self.mg._exec_cypher(
+ "MATCH (n {user_id: %s}) RETURN n", params=("u2",)
+ )
+ assert len(results) == 1
+
+ def test_reset_clears_everything(self):
+ self.mg._merge_node("u1", "alice", [0.0])
+ self.mg._merge_node("u2", "bob", [0.0])
+ self.mg.ag.commit()
+
+ self.mg.reset()
+
+ results = self.mg._exec_cypher("MATCH (n) RETURN n")
+ assert results == []
+
+
+# ==============================================================================
+# Test: _delete_entities
+# ==============================================================================
+
+@skip_no_age
+class TestDeleteEntities:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_deletes_specific_relationship(self):
+ self.mg._merge_node("u1", "alice", [0.0])
+ self.mg._merge_node("u1", "bob", [0.0])
+ self.mg.ag.commit()
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:LIKES]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg.ag.commit()
+
+ # Delete only KNOWS
+ self.mg._delete_entities(
+ [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}],
+ {"user_id": "u1"},
+ )
+
+ # LIKES should remain
+ results = self.mg._exec_cypher(
+ "MATCH (n {name: %s})-[r]->(m) RETURN n.name, type(r), m.name",
+ cols=["source", "rel", "target"],
+ params=("alice",),
+ )
+ assert len(results) == 1
+ assert results[0]["rel"] == "LIKES"
+
+
+# ==============================================================================
+# Test: _add_entities (full flow with merge + similarity)
+# ==============================================================================
+
+@skip_no_age
+class TestAddEntities:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_creates_new_nodes_and_relationship(self):
+ result = self.mg._add_entities(
+ [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}],
+ {"user_id": "u1"},
+ entity_type_map={"alice": "person", "bob": "person"},
+ )
+ assert len(result) == 1
+ assert result[0][0]["source"] == "alice"
+ assert result[0][0]["relationship"] == "KNOWS"
+ assert result[0][0]["target"] == "bob"
+
+ # Verify nodes exist in DB
+ nodes = self.mg._exec_cypher(
+ "MATCH (n {user_id: %s}) RETURN n", params=("u1",)
+ )
+ names = {n["name"] for n in nodes}
+ assert names == {"alice", "bob"}
+
+ def test_merges_to_existing_similar_node(self):
+ # Pre-create "alice" with a known embedding
+ emb = self.mg.embedding_model.embed("alice")
+ self.mg._merge_node("u1", "alice", emb)
+ self.mg.ag.commit()
+
+ # Now add an entity where source="alice" — should merge to existing
+ self.mg.threshold = 0.99 # high threshold, but same embedding = exact match
+ result = self.mg._add_entities(
+ [{"source": "alice", "destination": "carol", "relationship": "LIKES"}],
+ {"user_id": "u1"},
+ entity_type_map={"alice": "person", "carol": "person"},
+ )
+
+ # Should still have exactly one "alice" node (not a duplicate)
+ nodes = self.mg._exec_cypher(
+ "MATCH (n {user_id: %s, name: %s}) RETURN n",
+ params=("u1", "alice"),
+ )
+ assert len(nodes) == 1
+ # mentions should be > 1 from merge
+ assert nodes[0]["mentions"] >= 2
+
+ def test_add_multiple_relationships(self):
+ entities = [
+ {"source": "alice", "destination": "bob", "relationship": "KNOWS"},
+ {"source": "alice", "destination": "carol", "relationship": "LIKES"},
+ {"source": "bob", "destination": "carol", "relationship": "WORKS_WITH"},
+ ]
+ self.mg._add_entities(entities, {"user_id": "u1"}, entity_type_map={})
+
+ results = self.mg.get_all({"user_id": "u1"})
+ assert len(results) == 3
+ rels = {(r["source"], r["relationship"], r["target"]) for r in results}
+ assert ("alice", "KNOWS", "bob") in rels
+ assert ("alice", "LIKES", "carol") in rels
+ assert ("bob", "WORKS_WITH", "carol") in rels
+
+
+# ==============================================================================
+# Test: _search_graph_db
+# ==============================================================================
+
+@skip_no_age
+class TestSearchGraphDB:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_finds_related_entities(self):
+ # Create a small graph
+ emb_alice = self.mg.embedding_model.embed("alice")
+ emb_bob = self.mg.embedding_model.embed("bob")
+ self.mg._merge_node("u1", "alice", emb_alice)
+ self.mg._merge_node("u1", "bob", emb_bob)
+ self.mg.ag.commit()
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg.ag.commit()
+
+ # Search with "alice" embedding — should find the KNOWS relationship
+ self.mg.threshold = 0.99 # exact match only
+ results = self.mg._search_graph_db(["alice"], {"user_id": "u1"})
+ assert len(results) >= 1
+ found_knows = any(
+ r["source"] == "alice" and r["relationship"] == "KNOWS" and r["destination"] == "bob"
+ for r in results
+ )
+ assert found_knows, f"Expected KNOWS relationship in {results}"
+
+ def test_search_returns_empty_for_no_matches(self):
+ self.mg.threshold = 0.99
+ results = self.mg._search_graph_db(["nonexistent"], {"user_id": "u1"})
+ assert results == []
+
+
+# ==============================================================================
+# Test: Full add() + search() integration (mocking LLM, real DB)
+# ==============================================================================
+
+@skip_no_age
+class TestAddSearchIntegration:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_full_add_and_search_cycle(self):
+ """Simulates the full add() → search() cycle with mocked LLM responses."""
+ filters = {"user_id": "test_user_1"}
+
+ # Mock LLM: _retrieve_nodes_from_data
+ self.mg.llm.generate_response.side_effect = [
+ # 1st call: extract entities for add()
+ {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
+ {"entity": "Alice", "entity_type": "person"},
+ {"entity": "Bob", "entity_type": "person"},
+ ]}}]},
+ # 2nd call: establish relations for add()
+ {"tool_calls": [{"name": "add_entities", "arguments": {"entities": [
+ {"source": "Alice", "relationship": "knows", "destination": "Bob"},
+ ]}}]},
+ # 3rd call: get_delete_entities (nothing to delete)
+ {"tool_calls": []},
+ # 4th call: extract entities for search()
+ {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
+ {"entity": "Alice", "entity_type": "person"},
+ ]}}]},
+ ]
+
+ # Add
+ add_result = self.mg.add("Alice knows Bob", filters)
+ assert "added_entities" in add_result
+ assert "deleted_entities" in add_result
+
+ # Verify in DB
+ all_rels = self.mg.get_all(filters)
+ assert len(all_rels) == 1
+ assert all_rels[0]["source"] == "alice"
+ # Relationship labels are lowercased by _remove_spaces_from_entities
+ assert all_rels[0]["relationship"] == "knows"
+ assert all_rels[0]["target"] == "bob"
+
+ # Search
+ search_results = self.mg.search("Who does Alice know?", filters)
+ assert len(search_results) >= 1
+ assert any(r["source"] == "alice" and r["destination"] == "bob" for r in search_results)
+
+ # Delete all
+ self.mg.delete_all(filters)
+ remaining = self.mg.get_all(filters)
+ assert remaining == []
+
+ def test_add_then_update_relationship(self):
+ """Tests that adding conflicting data removes old relationships."""
+ filters = {"user_id": "test_user_2"}
+
+ # First add: Alice likes cats
+ self.mg.llm.generate_response.side_effect = [
+ {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
+ {"entity": "Alice", "entity_type": "person"},
+ {"entity": "cats", "entity_type": "animal"},
+ ]}}]},
+ {"tool_calls": [{"name": "add_entities", "arguments": {"entities": [
+ {"source": "Alice", "relationship": "likes", "destination": "cats"},
+ ]}}]},
+ {"tool_calls": []}, # nothing to delete
+ ]
+ self.mg.add("Alice likes cats", filters)
+
+ all_rels = self.mg.get_all(filters)
+ assert len(all_rels) == 1
+ assert all_rels[0]["relationship"] == "likes"
+
+ # Second add: Alice now dislikes cats (delete old, add new)
+ self.mg.llm.generate_response.side_effect = [
+ {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
+ {"entity": "Alice", "entity_type": "person"},
+ {"entity": "cats", "entity_type": "animal"},
+ ]}}]},
+ {"tool_calls": [{"name": "add_entities", "arguments": {"entities": [
+ {"source": "Alice", "relationship": "dislikes", "destination": "cats"},
+ ]}}]},
+ # LLM says to delete the old likes relationship
+ {"tool_calls": [{"name": "delete_graph_memory", "arguments": {
+ "source": "alice", "relationship": "likes", "destination": "cats",
+ }}]},
+ ]
+ self.mg.add("Alice dislikes cats", filters)
+
+ all_rels = self.mg.get_all(filters)
+ rels = {r["relationship"] for r in all_rels}
+ assert "likes" not in rels
+ assert "dislikes" in rels
+
+ self.mg.delete_all(filters)
+
+
+# ==============================================================================
+# Test: Multi-tenant isolation
+# ==============================================================================
+
+@skip_no_age
+class TestMultiTenantIsolation:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_users_cant_see_each_others_data(self):
+ # User 1
+ self.mg._merge_node("user_1", "alice", [0.0]*16)
+ self.mg._merge_node("user_1", "bob", [0.0]*16)
+ self.mg.ag.commit()
+ self.mg._exec_cypher(
+ "MATCH (s {user_id: %s, name: %s}), (d {user_id: %s, name: %s}) "
+ "MERGE (s)-[:KNOWS]->(d)",
+ params=("user_1", "alice", "user_1", "bob"),
+ )
+ self.mg.ag.commit()
+
+ # User 2
+ self.mg._merge_node("user_2", "carol", [0.0]*16)
+ self.mg._merge_node("user_2", "dave", [0.0]*16)
+ self.mg.ag.commit()
+ self.mg._exec_cypher(
+ "MATCH (s {user_id: %s, name: %s}), (d {user_id: %s, name: %s}) "
+ "MERGE (s)-[:WORKS_WITH]->(d)",
+ params=("user_2", "carol", "user_2", "dave"),
+ )
+ self.mg.ag.commit()
+
+ # User 1 sees only their data
+ u1_results = self.mg.get_all({"user_id": "user_1"})
+ assert len(u1_results) == 1
+ assert u1_results[0]["source"] == "alice"
+
+ # User 2 sees only their data
+ u2_results = self.mg.get_all({"user_id": "user_2"})
+ assert len(u2_results) == 1
+ assert u2_results[0]["source"] == "carol"
+
+ # Delete user 1 doesn't affect user 2
+ self.mg.delete_all({"user_id": "user_1"})
+ u2_after = self.mg.get_all({"user_id": "user_2"})
+ assert len(u2_after) == 1
+
+
+# ==============================================================================
+# Test: agent_id / run_id filtering in delete_all and get_all
+# ==============================================================================
+
+@skip_no_age
+class TestAgentRunIdFiltering:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_get_all_filters_by_agent_id(self):
+ # Create nodes for two different agents under the same user
+ self.mg._merge_node("u1", "alice", [0.0]*16, agent_id="agent_a")
+ self.mg._merge_node("u1", "bob", [0.0]*16, agent_id="agent_a")
+ self.mg._merge_node("u1", "carol", [0.0]*16, agent_id="agent_b")
+ self.mg._merge_node("u1", "dave", [0.0]*16, agent_id="agent_b")
+ self.mg.ag.commit()
+
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:LIKES]->(d)",
+ params=("carol", "dave"),
+ )
+ self.mg.ag.commit()
+
+ # get_all with agent_a should only return alice->bob
+ results_a = self.mg.get_all({"user_id": "u1", "agent_id": "agent_a"})
+ assert len(results_a) == 1
+ assert results_a[0]["source"] == "alice"
+
+ # get_all with agent_b should only return carol->dave
+ results_b = self.mg.get_all({"user_id": "u1", "agent_id": "agent_b"})
+ assert len(results_b) == 1
+ assert results_b[0]["source"] == "carol"
+
+ def test_delete_all_filters_by_agent_id(self):
+ self.mg._merge_node("u1", "alice", [0.0]*16, agent_id="agent_a")
+ self.mg._merge_node("u1", "bob", [0.0]*16, agent_id="agent_a")
+ self.mg._merge_node("u1", "carol", [0.0]*16, agent_id="agent_b")
+ self.mg.ag.commit()
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:REL]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg.ag.commit()
+
+ # Delete only agent_a's data
+ self.mg.delete_all({"user_id": "u1", "agent_id": "agent_a"})
+
+ # agent_a nodes should be gone
+ nodes_a = self.mg._exec_cypher(
+ "MATCH (n) WHERE n.user_id = %s AND n.agent_id = %s RETURN n",
+ params=("u1", "agent_a"),
+ )
+ assert nodes_a == []
+
+ # agent_b's data should still exist
+ nodes_b = self.mg._exec_cypher(
+ "MATCH (n) WHERE n.user_id = %s AND n.agent_id = %s RETURN n",
+ params=("u1", "agent_b"),
+ )
+ assert len(nodes_b) == 1
+
+ def test_get_all_filters_by_run_id(self):
+ self.mg._merge_node("u1", "alice", [0.0]*16, run_id="run_1")
+ self.mg._merge_node("u1", "bob", [0.0]*16, run_id="run_1")
+ self.mg._merge_node("u1", "carol", [0.0]*16, run_id="run_2")
+ self.mg._merge_node("u1", "dave", [0.0]*16, run_id="run_2")
+ self.mg.ag.commit()
+
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:R1]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:R2]->(d)",
+ params=("carol", "dave"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg.get_all({"user_id": "u1", "run_id": "run_1"})
+ assert len(results) == 1
+ assert results[0]["source"] == "alice"
+
+ def test_delete_all_filters_by_run_id(self):
+ self.mg._merge_node("u1", "alice", [0.0]*16, run_id="run_1")
+ self.mg._merge_node("u1", "bob", [0.0]*16, run_id="run_2")
+ self.mg.ag.commit()
+
+ self.mg.delete_all({"user_id": "u1", "run_id": "run_1"})
+
+ # run_1 node should be gone
+ nodes_1 = self.mg._exec_cypher(
+ "MATCH (n) WHERE n.user_id = %s AND n.run_id = %s RETURN n",
+ params=("u1", "run_1"),
+ )
+ assert nodes_1 == []
+
+ # run_2 node should remain
+ nodes_2 = self.mg._exec_cypher(
+ "MATCH (n) WHERE n.user_id = %s AND n.run_id = %s RETURN n",
+ params=("u1", "run_2"),
+ )
+ assert len(nodes_2) == 1
+
+
+# ==============================================================================
+# Test: Special characters and edge cases
+# ==============================================================================
+
+@skip_no_age
+class TestEdgeCases:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_node_name_with_underscores(self):
+ """Entity names go through _remove_spaces_from_entities which lowercases
+ and replaces spaces with underscores."""
+ self.mg._merge_node("u1", "new_york_city", [0.0]*16)
+ self.mg._merge_node("u1", "united_states", [0.0]*16)
+ self.mg.ag.commit()
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:LOCATED_IN]->(d)",
+ params=("new_york_city", "united_states"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg.get_all({"user_id": "u1"})
+ assert len(results) == 1
+ assert results[0]["source"] == "new_york_city"
+ assert results[0]["target"] == "united_states"
+
+ def test_node_with_apostrophe_in_name(self):
+ """The AGE Python driver has a known limitation where single quotes in
+ parameterized values cause a syntax error due to double-quoting in
+ buildCypher(). In practice this is not hit because entity names go
+ through _remove_spaces_from_entities which sanitizes them."""
+ import psycopg2
+ with pytest.raises(psycopg2.errors.SyntaxError):
+ self.mg._merge_node("u1", "o'brien", [0.0]*16)
+
+ def test_empty_graph_get_all(self):
+ results = self.mg.get_all({"user_id": "u1"})
+ assert results == []
+
+ def test_empty_graph_delete_all_no_error(self):
+ # Should not raise even on empty graph
+ self.mg.delete_all({"user_id": "u1"})
+
+ def test_empty_graph_reset_no_error(self):
+ self.mg.reset()
+
+ def test_duplicate_relationship_merge_is_idempotent(self):
+ """MERGE on the same relationship twice should not create duplicates."""
+ self.mg._merge_node("u1", "alice", [0.0]*16)
+ self.mg._merge_node("u1", "bob", [0.0]*16)
+ self.mg.ag.commit()
+
+ for _ in range(3):
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:FRIENDS]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg.get_all({"user_id": "u1"})
+ assert len(results) == 1 # Only one relationship, not 3
+
+ def test_bidirectional_relationships(self):
+ """Two nodes can have relationships in both directions."""
+ self.mg._merge_node("u1", "alice", [0.0]*16)
+ self.mg._merge_node("u1", "bob", [0.0]*16)
+ self.mg.ag.commit()
+
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:FOLLOWS]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:FOLLOWS]->(d)",
+ params=("bob", "alice"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg.get_all({"user_id": "u1"})
+ assert len(results) == 2
+ pairs = {(r["source"], r["target"]) for r in results}
+ assert ("alice", "bob") in pairs
+ assert ("bob", "alice") in pairs
+
+ def test_self_referencing_relationship(self):
+ """A node can have a relationship to itself."""
+ self.mg._merge_node("u1", "alice", [0.0]*16)
+ self.mg.ag.commit()
+
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS_SELF]->(d)",
+ params=("alice", "alice"),
+ )
+ self.mg.ag.commit()
+
+ results = self.mg.get_all({"user_id": "u1"})
+ assert len(results) == 1
+ assert results[0]["source"] == "alice"
+ assert results[0]["target"] == "alice"
+
+
+# ==============================================================================
+# Test: _search_graph_db with agent_id/run_id filtering
+# ==============================================================================
+
+@skip_no_age
+class TestSearchGraphDBFiltering:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_search_filters_by_agent_id(self):
+ emb = self.mg.embedding_model.embed("alice")
+ self.mg._merge_node("u1", "alice", emb, agent_id="agent_a")
+ self.mg._merge_node("u1", "bob", [0.0]*16, agent_id="agent_a")
+ self.mg._merge_node("u1", "alice_clone", emb, agent_id="agent_b")
+ self.mg._merge_node("u1", "carol", [0.0]*16, agent_id="agent_b")
+ self.mg.ag.commit()
+
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:R1]->(d)",
+ params=("alice", "bob"),
+ )
+ self.mg._exec_cypher(
+ "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:R2]->(d)",
+ params=("alice_clone", "carol"),
+ )
+ self.mg.ag.commit()
+
+ self.mg.threshold = 0.99
+ results = self.mg._search_graph_db(
+ ["alice"], {"user_id": "u1", "agent_id": "agent_a"}
+ )
+ # Should only find relationships for agent_a
+ for r in results:
+ assert r["source"] != "alice_clone", f"Leaked agent_b data: {r}"
+
+ def test_find_similar_node_filters_by_run_id(self):
+ emb = [1.0, 0.0, 0.0, 0.0] + [0.0]*12
+ self.mg._merge_node("u1", "target_run1", emb, run_id="run_1")
+ self.mg._merge_node("u1", "target_run2", emb, run_id="run_2")
+ self.mg.ag.commit()
+
+ match = self.mg._find_similar_node(
+ emb, {"user_id": "u1", "run_id": "run_1"}, threshold=0.99
+ )
+ assert match is not None
+ assert match["name"] == "target_run1"
+
+
+# ==============================================================================
+# Test: Full lifecycle with agent_id
+# ==============================================================================
+
+@skip_no_age
+class TestFullLifecycleWithAgentId:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_add_search_delete_with_agent_id(self):
+ """Full cycle: add → get_all → search → delete_all, all scoped by agent_id."""
+ filters = {"user_id": "u1", "agent_id": "agent_x"}
+
+ self.mg.llm.generate_response.side_effect = [
+ # extract entities
+ {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
+ {"entity": "Python", "entity_type": "language"},
+ {"entity": "Alice", "entity_type": "person"},
+ ]}}]},
+ # establish relations
+ {"tool_calls": [{"name": "add_entities", "arguments": {"entities": [
+ {"source": "Alice", "relationship": "uses", "destination": "Python"},
+ ]}}]},
+ # nothing to delete
+ {"tool_calls": []},
+ ]
+
+ self.mg.add("Alice uses Python", filters)
+
+ # Verify nodes have agent_id
+ nodes = self.mg._exec_cypher(
+ "MATCH (n) WHERE n.user_id = %s AND n.agent_id = %s RETURN n",
+ params=("u1", "agent_x"),
+ )
+ assert len(nodes) == 2
+ names = {n["name"] for n in nodes}
+ assert names == {"alice", "python"}
+
+ # get_all with agent_id filter
+ all_rels = self.mg.get_all(filters)
+ assert len(all_rels) == 1
+ assert all_rels[0]["relationship"] == "uses"
+
+ # search
+ self.mg.llm.generate_response.side_effect = [
+ {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
+ {"entity": "Alice", "entity_type": "person"},
+ ]}}]},
+ ]
+ search_results = self.mg.search("What does Alice use?", filters)
+ assert len(search_results) >= 1
+
+ # delete only agent_x
+ self.mg.delete_all(filters)
+ remaining = self.mg.get_all(filters)
+ assert remaining == []
+
+
+# ==============================================================================
+# Test: _merge_node preserves created timestamp
+# ==============================================================================
+
+@skip_no_age
+class TestMergeNodeTimestamp:
+
+ def setup_method(self):
+ self.mg = _make_e2e_instance()
+
+ def teardown_method(self):
+ _cleanup(self.mg)
+
+ def test_created_preserved_across_merges(self):
+ """The created timestamp should be set on first merge and preserved on subsequent merges."""
+ self.mg._merge_node("u1", "alice", [0.0]*16)
+ self.mg.ag.commit()
+
+ nodes1 = self.mg._exec_cypher(
+ "MATCH (n {name: %s}) RETURN n", params=("alice",)
+ )
+ created1 = nodes1[0]["created"]
+ assert created1 is not None
+
+ # Second merge — created should not change
+ import time
+ time.sleep(0.05) # Ensure clock moves
+ self.mg._merge_node("u1", "alice", [0.0]*16)
+ self.mg.ag.commit()
+
+ nodes2 = self.mg._exec_cypher(
+ "MATCH (n {name: %s}) RETURN n", params=("alice",)
+ )
+ created2 = nodes2[0]["created"]
+ assert created2 == created1, f"created changed from {created1} to {created2}"
+ assert nodes2[0]["mentions"] == 2
diff --git a/tests/memory/test_apache_age_memory.py b/tests/memory/test_apache_age_memory.py
new file mode 100644
index 000000000..a0d61b39f
--- /dev/null
+++ b/tests/memory/test_apache_age_memory.py
@@ -0,0 +1,221 @@
+from unittest.mock import MagicMock, Mock, patch
+
+# age and rank_bm25 are optional deps — mock them so tests run without install
+_age_mock = Mock()
+patch.dict("sys.modules", {
+ "age": _age_mock,
+ "age.models": Mock(),
+ "rank_bm25": Mock(),
+}).start()
+
+from mem0.memory.apache_age_memory import MemoryGraph, _cosine_similarity # noqa: E402
+
+
+def _make_instance():
+ with patch.object(MemoryGraph, "__init__", return_value=None):
+ instance = MemoryGraph.__new__(MemoryGraph)
+ instance.llm_provider = "openai"
+ instance.llm = MagicMock()
+ instance.embedding_model = MagicMock()
+ instance.config = MagicMock()
+ instance.config.graph_store.custom_prompt = None
+ instance.ag = MagicMock()
+ instance.graph_name = "test_graph"
+ instance.threshold = 0.7
+ return instance
+
+
+class TestCosineSimilarity:
+ """Tests for the _cosine_similarity helper."""
+
+ def test_identical_vectors(self):
+ assert abs(_cosine_similarity([1, 0, 0], [1, 0, 0]) - 1.0) < 1e-6
+
+ def test_orthogonal_vectors(self):
+ assert abs(_cosine_similarity([1, 0, 0], [0, 1, 0])) < 1e-6
+
+ def test_zero_vector(self):
+ assert _cosine_similarity([0, 0, 0], [1, 2, 3]) == 0.0
+
+
+class TestRetrieveNodesFromData:
+ """Tests for _retrieve_nodes_from_data in Apache AGE MemoryGraph."""
+
+ def test_normal_entities_extracted(self):
+ instance = _make_instance()
+ instance.llm.generate_response.return_value = {
+ "tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
+ {"entity": "Alice", "entity_type": "person"},
+ {"entity": "hiking", "entity_type": "activity"},
+ ]}}]
+ }
+ result = instance._retrieve_nodes_from_data("Alice loves hiking", {"user_id": "u1"})
+ assert result == {"alice": "person", "hiking": "activity"}
+
+ def test_malformed_entity_missing_entity_type_is_skipped(self):
+ instance = _make_instance()
+ instance.llm.generate_response.return_value = {
+ "tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
+ {"entity": "matrix multiplication", "entity_type": "task"},
+ {"entity": "task"},
+ {"entity": "ReLU", "entity_type": "task"},
+ ]}}]
+ }
+ result = instance._retrieve_nodes_from_data("some text", {"user_id": "u1"})
+ assert "matrix_multiplication" in result
+ assert "relu" in result
+ assert "task" not in result
+
+ def test_missing_entities_key_returns_empty(self):
+ instance = _make_instance()
+ instance.llm.generate_response.return_value = {
+ "tool_calls": [{"name": "extract_entities", "arguments": {"text": "Hello."}}]
+ }
+ result = instance._retrieve_nodes_from_data("Hello.", {"user_id": "u1"})
+ assert result == {}
+
+ def test_none_tool_calls_returns_empty(self):
+ instance = _make_instance()
+ instance.llm.generate_response.return_value = {"tool_calls": None}
+ result = instance._retrieve_nodes_from_data("hello world", {"user_id": "u1"})
+ assert result == {}
+
+
+class TestEstablishNodesRelationsFromData:
+ """Tests for _establish_nodes_relations_from_data in Apache AGE MemoryGraph."""
+
+ def test_none_response_does_not_crash(self):
+ instance = _make_instance()
+ instance.llm.generate_response.return_value = None
+ result = instance._establish_nodes_relations_from_data(
+ "Hello world", {"user_id": "u1"}, {}
+ )
+ assert result == []
+
+ def test_empty_tool_calls_returns_empty(self):
+ instance = _make_instance()
+ instance.llm.generate_response.return_value = {"tool_calls": []}
+ result = instance._establish_nodes_relations_from_data(
+ "Hello world", {"user_id": "u1"}, {}
+ )
+ assert result == []
+
+ def test_valid_entities_returned(self):
+ instance = _make_instance()
+ instance.llm.generate_response.return_value = {
+ "tool_calls": [{"name": "add_entities", "arguments": {"entities": [
+ {"source": "alice", "relationship": "loves", "destination": "hiking"}
+ ]}}]
+ }
+ result = instance._establish_nodes_relations_from_data(
+ "Alice loves hiking", {"user_id": "u1"}, {"alice": "person"}
+ )
+ assert len(result) == 1
+ assert result[0]["source"] == "alice"
+
+
+class TestRemoveSpacesFromEntities:
+ """Tests for _remove_spaces_from_entities."""
+
+ def test_spaces_and_case(self):
+ instance = _make_instance()
+ entities = [{"source": "Alice Smith", "relationship": "Works At", "destination": "Big Corp"}]
+ result = instance._remove_spaces_from_entities(entities)
+ assert result[0]["source"] == "alice_smith"
+ assert result[0]["relationship"] == "works_at"
+ assert result[0]["destination"] == "big_corp"
+
+
+class TestFindSimilarNode:
+ """Tests for _find_similar_node."""
+
+ def test_returns_none_when_no_nodes(self):
+ instance = _make_instance()
+ instance._exec_cypher = MagicMock(return_value=[])
+ result = instance._find_similar_node([1.0, 0.0], {"user_id": "u1"}, threshold=0.9)
+ assert result is None
+
+ def test_returns_best_match_above_threshold(self):
+ instance = _make_instance()
+ instance._exec_cypher = MagicMock(return_value=[
+ {"name": "alice", "embedding": [1.0, 0.0], "user_id": "u1"},
+ {"name": "bob", "embedding": [0.0, 1.0], "user_id": "u1"},
+ ])
+ result = instance._find_similar_node([1.0, 0.0], {"user_id": "u1"}, threshold=0.9)
+ assert result["name"] == "alice"
+
+ def test_filters_by_agent_id(self):
+ instance = _make_instance()
+ instance._exec_cypher = MagicMock(return_value=[
+ {"name": "alice", "embedding": [1.0, 0.0], "user_id": "u1", "agent_id": "a2"},
+ ])
+ result = instance._find_similar_node(
+ [1.0, 0.0], {"user_id": "u1", "agent_id": "a1"}, threshold=0.9
+ )
+ assert result is None
+
+
+class TestDeleteAll:
+ """Tests for delete_all."""
+
+ def test_calls_exec_cypher_and_commits(self):
+ instance = _make_instance()
+ instance._exec_cypher = MagicMock(return_value=[])
+ instance.delete_all({"user_id": "u1"})
+ instance._exec_cypher.assert_called_once()
+ instance.ag.commit.assert_called_once()
+
+
+class TestGetAll:
+ """Tests for get_all."""
+
+ def test_returns_formatted_results(self):
+ instance = _make_instance()
+ instance._exec_cypher = MagicMock(return_value=[
+ {"source": "alice", "relationship": "KNOWS", "target": "bob"},
+ {"source": "alice", "relationship": "LIKES", "target": "hiking"},
+ ])
+ results = instance.get_all({"user_id": "u1"}, limit=10)
+ assert len(results) == 2
+ assert results[0]["source"] == "alice"
+ assert results[0]["relationship"] == "KNOWS"
+ assert results[0]["target"] == "bob"
+
+ def test_respects_limit(self):
+ instance = _make_instance()
+ instance._exec_cypher = MagicMock(return_value=[
+ {"source": f"n{i}", "relationship": "R", "target": f"m{i}"} for i in range(10)
+ ])
+ results = instance.get_all({"user_id": "u1"}, limit=3)
+ assert len(results) == 3
+
+
+class TestAdd:
+ """Tests for the add orchestration method."""
+
+ def test_add_returns_added_and_deleted(self):
+ instance = _make_instance()
+ instance._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"})
+ instance._establish_nodes_relations_from_data = MagicMock(return_value=[
+ {"source": "alice", "relationship": "knows", "destination": "bob"}
+ ])
+ instance._search_graph_db = MagicMock(return_value=[])
+ instance._get_delete_entities_from_search_output = MagicMock(return_value=[])
+ instance._delete_entities = MagicMock(return_value=[])
+ instance._add_entities = MagicMock(return_value=["added"])
+
+ result = instance.add("Alice knows Bob", {"user_id": "u1"})
+ assert "deleted_entities" in result
+ assert "added_entities" in result
+ assert result["added_entities"] == ["added"]
+
+
+class TestSearch:
+ """Tests for the search method."""
+
+ def test_returns_empty_when_no_search_output(self):
+ instance = _make_instance()
+ instance._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"})
+ instance._search_graph_db = MagicMock(return_value=[])
+ result = instance.search("Who is Alice?", {"user_id": "u1"})
+ assert result == []