From 305ce7b6b3b3feec49f00011b2a1ea172b38d8df Mon Sep 17 00:00:00 2001 From: Utkarsh Date: Fri, 20 Mar 2026 20:27:48 +0530 Subject: [PATCH] feat: add Apache AGE graph store support (#4448) Co-authored-by: utkarsh240799 Co-authored-by: Claude Opus 4.6 (1M context) --- docs/open-source/features/graph-memory.mdx | 40 +- mem0/graphs/configs.py | 26 +- mem0/memory/apache_age_memory.py | 573 ++++++++++ mem0/utils/factory.py | 1 + pyproject.toml | 1 + tests/memory/test_apache_age_e2e.py | 1105 ++++++++++++++++++++ tests/memory/test_apache_age_memory.py | 226 ++++ 7 files changed, 1968 insertions(+), 4 deletions(-) create mode 100644 mem0/memory/apache_age_memory.py create mode 100644 tests/memory/test_apache_age_e2e.py create mode 100644 tests/memory/test_apache_age_memory.py 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..74c51b555 --- /dev/null +++ b/mem0/memory/apache_age_memory.py @@ -0,0 +1,573 @@ +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) + + +def _get_similar_nodes(nodes, query_embedding, filters, threshold): + """Find nodes above the similarity threshold from a fetched node list. + + Shared by ``_find_similar_node`` and ``_search_graph_db`` to avoid + duplicating the client-side cosine similarity logic. + """ + matches = [] + 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(query_embedding, stored_emb) + if sim >= threshold: + matches.append({"name": props.get("name"), "similarity": sim, "props": props}) + + matches.sort(key=lambda x: x["similarity"], reverse=True) + return matches + + +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 _fetch_user_nodes_with_embeddings(self, user_id): + """Fetch all nodes with embeddings for a given user_id.""" + return self._exec_cypher( + "MATCH (n {user_id: %s}) WHERE n.embedding IS NOT NULL RETURN n", + params=(user_id,), + ) + + 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._fetch_user_nodes_with_embeddings(filters["user_id"]) + matches = _get_similar_nodes(nodes, embedding, filters, threshold) + return matches[0]["props"] if matches else None + + 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 + and optional filter properties are set in a single query. + """ + set_parts = [ + "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: + set_parts.append("n.agent_id = %s") + params.append(agent_id) + if run_id: + set_parts.append("n.run_id = %s") + params.append(run_id) + + set_clause = ", ".join(set_parts) + self._exec_cypher( + f"MERGE (n {{user_id: %s, name: %s}}) SET {set_clause}", + params=tuple(params), + ) + + def close(self): + """Close the underlying database connection.""" + if self.ag: + self.ag.close() + + # -- 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: + list: A list of dicts with keys "source", "relationship", "destination". + """ + 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.""" + 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) + params.append(limit) + + results = self._exec_cypher( + f"MATCH (n)-[r]->(m) WHERE {where_clause} " + f"RETURN n.name, type(r), m.name LIMIT %s", + cols=["source", "relationship", "target"], + params=tuple(params), + ) + + final_results = [] + for result in results: + 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) + + nodes = self._fetch_user_nodes_with_embeddings(filters["user_id"]) + similar_nodes = _get_similar_nodes(nodes, n_embedding, filters, self.threshold) + + # 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 = [] + + try: + 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() + except Exception: + self.ag.rollback() + raise + + 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 = [] + + try: + 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), + ) + results.append(result) + + self.ag.commit() + except Exception: + self.ag.rollback() + raise + + 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 afbd8263f..87a76c885 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -218,6 +218,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..5919136d6 --- /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 # noqa: E402 + +# -- 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 + 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..e0823cb4d --- /dev/null +++ b/tests/memory/test_apache_age_memory.py @@ -0,0 +1,226 @@ +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_passes_limit_to_cypher(self): + """Limit is enforced via LIMIT in the Cypher query, not Python slicing.""" + instance = _make_instance() + instance._exec_cypher = MagicMock(return_value=[ + {"source": "n0", "relationship": "R", "target": "m0"}, + ]) + instance.get_all({"user_id": "u1"}, limit=3) + # Verify limit was passed as a parameter to the query + cypher_stmt = instance._exec_cypher.call_args[0][0] + assert "LIMIT %s" in cypher_stmt + params = instance._exec_cypher.call_args[1].get("params") or instance._exec_cypher.call_args[0][2] + assert 3 in params + + +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 == []