diff --git a/examples/basic_modules/neo4j_example.py b/examples/basic_modules/neo4j_example.py index bf7bdf9c5..082ad8c3e 100644 --- a/examples/basic_modules/neo4j_example.py +++ b/examples/basic_modules/neo4j_example.py @@ -42,7 +42,6 @@ def example_multi_db(db_name: str = "paper"): metadata=TreeNodeTextualMemoryMetadata( memory_type="LongTermMemory", key="Multi-UAV Long-Term Coverage", - value="Research topic on distributed multi-agent UAV navigation and coverage", hierarchy_level="topic", type="fact", memory_time="2024-01-01", @@ -74,7 +73,6 @@ def example_multi_db(db_name: str = "paper"): metadata=TreeNodeTextualMemoryMetadata( memory_type="LongTermMemory", key="Reward Function Design", - value="Combines coverage, energy efficiency, and overlap penalty", hierarchy_level="concept", type="fact", memory_time="2024-01-01", @@ -99,7 +97,6 @@ def example_multi_db(db_name: str = "paper"): metadata=TreeNodeTextualMemoryMetadata( memory_type="LongTermMemory", key="Energy Model", - value="Includes communication and motion energy consumption", hierarchy_level="concept", type="fact", memory_time="2024-01-01", @@ -122,7 +119,6 @@ def example_multi_db(db_name: str = "paper"): metadata=TreeNodeTextualMemoryMetadata( memory_type="LongTermMemory", key="Coverage Metrics", - value="CT and FT used for long-term area and fairness evaluation", hierarchy_level="concept", type="fact", memory_time="2024-01-01", @@ -161,7 +157,6 @@ def example_multi_db(db_name: str = "paper"): metadata=TreeNodeTextualMemoryMetadata( memory_type="WorkingMemory", key="Reward Components", - value="Coverage gain, energy usage penalty, overlap penalty", hierarchy_level="fact", type="fact", memory_time="2024-01-01", @@ -186,7 +181,6 @@ def example_multi_db(db_name: str = "paper"): metadata=TreeNodeTextualMemoryMetadata( memory_type="LongTermMemory", key="Energy Cost Components", - value="Includes movement and communication energy", hierarchy_level="fact", type="fact", memory_time="2024-01-01", @@ -211,7 +205,6 @@ def example_multi_db(db_name: str = "paper"): metadata=TreeNodeTextualMemoryMetadata( memory_type="LongTermMemory", key="CT and FT Definition", - value="CT: total coverage duration; FT: fairness index", hierarchy_level="fact", type="fact", memory_time="2024-01-01", @@ -347,9 +340,202 @@ def example_shared_db(db_name: str = "shared-traval-group"): print(graph_alice.get_node(node["id"])) +def run_user_session( + user_name: str, + db_name: str, + topic_text: str, + concept_texts: list[str], + fact_texts: list[str], + community: bool = False, +): + print(f"\n=== {user_name} starts building their memory graph ===") + + # Manually initialize correct GraphDB class + if community: + config = GraphDBConfigFactory( + backend="neo4j-community", + config={ + "uri": "bolt://localhost:7687", + "user": "neo4j", + "password": "12345678", + "db_name": db_name, + "user_name": user_name, + "use_multi_db": False, + "auto_create": False, # Neo4j Community does not allow auto DB creation + "embedding_dimension": 768, + "vec_config": { + # Pass nested config to initialize external vector DB + # If you use qdrant, please use Server instead of local mode. + "backend": "qdrant", + "config": { + "collection_name": "neo4j_vec_db", + "vector_dimension": 768, + "distance_metric": "cosine", + "host": "localhost", + "port": 6333, + }, + }, + }, + ) + else: + config = GraphDBConfigFactory( + backend="neo4j", + config={ + "uri": "bolt://localhost:7687", + "user": "neo4j", + "password": "12345678", + "db_name": db_name, + "user_name": user_name, + "use_multi_db": False, + "auto_create": True, + "embedding_dimension": 768, + }, + ) + graph = GraphStoreFactory.from_config(config) + + # Start with a clean slate for this user + graph.clear() + + now = datetime.utcnow().isoformat() + + # === Step 1: Create a root topic node (e.g., user's research focus) === + topic = TextualMemoryItem( + memory=topic_text, + metadata=TreeNodeTextualMemoryMetadata( + memory_type="LongTermMemory", + key="Research Topic", + hierarchy_level="topic", + type="fact", + memory_time="2024-01-01", + status="activated", + visibility="public", + updated_at=now, + embedding=embed_memory_item(topic_text), + ), + ) + graph.add_node(topic.id, topic.memory, topic.metadata.model_dump(exclude_none=True)) + + # === Step 2: Create two concept nodes linked to the topic === + concept_items = [] + for i, text in enumerate(concept_texts): + concept = TextualMemoryItem( + memory=text, + metadata=TreeNodeTextualMemoryMetadata( + memory_type="LongTermMemory", + key=f"Concept {i + 1}", + hierarchy_level="concept", + type="fact", + memory_time="2024-01-01", + status="activated", + visibility="public", + updated_at=now, + embedding=embed_memory_item(text), + tags=["concept"], + confidence=90 + i, + ), + ) + graph.add_node(concept.id, concept.memory, concept.metadata.model_dump(exclude_none=True)) + graph.add_edge(topic.id, concept.id, type="PARENT") + concept_items.append(concept) + + # === Step 3: Create supporting facts under each concept === + for i, text in enumerate(fact_texts): + fact = TextualMemoryItem( + memory=text, + metadata=TreeNodeTextualMemoryMetadata( + memory_type="WorkingMemory", + key=f"Fact {i + 1}", + hierarchy_level="fact", + type="fact", + memory_time="2024-01-01", + status="activated", + visibility="public", + updated_at=now, + embedding=embed_memory_item(text), + confidence=85.0, + tags=["fact"], + ), + ) + graph.add_node(fact.id, fact.memory, fact.metadata.model_dump(exclude_none=True)) + graph.add_edge(concept_items[i % len(concept_items)].id, fact.id, type="PARENT") + + # === Step 4: Retrieve memory using semantic search === + vector = embed_memory_item("How is memory retrieved?") + search_result = graph.search_by_embedding(vector, top_k=2) + for r in search_result: + node = graph.get_node(r["id"]) + print("πŸ” Search result:", node["memory"]) + + # === Step 5: Tag-based neighborhood discovery === + neighbors = graph.get_neighbors_by_tag(["concept"], exclude_ids=[], top_k=2) + print("πŸ“Ž Tag-related nodes:", [neighbor["memory"] for neighbor in neighbors]) + + # === Step 6: Retrieve children (facts) of first concept === + children = graph.get_children_with_embeddings(concept_items[0].id) + print("πŸ“ Children of concept:", [child["memory"] for child in children]) + + # === Step 7: Export a local subgraph and grouped statistics === + subgraph = graph.get_subgraph(topic.id, depth=2) + print("πŸ“Œ Subgraph node count:", len(subgraph["neighbors"])) + + stats = graph.get_grouped_counts(["memory_type", "status"]) + print("πŸ“Š Grouped counts:", stats) + + # === Step 8: Demonstrate updates and cleanup === + graph.update_node(concept_items[0].id, {"confidence": 99.0}) + graph.remove_oldest_memory("WorkingMemory", keep_latest=1) + graph.delete_edge(topic.id, concept_items[0].id, type="PARENT") + graph.delete_node(concept_items[1].id) + + # === Step 9: Export and re-import the entire graph structure === + exported = graph.export_graph() + graph.import_graph(exported) + print("πŸ“¦ Graph exported and re-imported, total nodes:", len(exported["nodes"])) + + +def example_complex_shared_db(db_name: str = "shared-traval-group-complex", community=False): + # User 1: Alice explores structured memory for LLMs + run_user_session( + user_name="alice", + db_name=db_name, + topic_text="Alice studies structured memory and long-term memory optimization in LLMs.", + concept_texts=[ + "Short-term memory can be simulated using WorkingMemory blocks.", + "A structured memory graph improves retrieval precision for agents.", + ], + fact_texts=[ + "Embedding search is used to find semantically similar memory items.", + "User memories are stored as node-edge structures that support hierarchical reasoning.", + ], + community=community, + ) + + # User 2: Bob focuses on GNN-based reasoning + run_user_session( + user_name="bob", + db_name=db_name, + topic_text="Bob investigates how graph neural networks can support knowledge reasoning.", + concept_texts=[ + "GNNs can learn high-order relations among entities.", + "Attention mechanisms in graphs improve inference precision.", + ], + fact_texts=[ + "GAT outperforms GCN in graph classification tasks.", + "Multi-hop reasoning helps answer complex queries.", + ], + community=community, + ) + + if __name__ == "__main__": print("\n=== Example: Multi-DB ===") example_multi_db(db_name="paper") print("\n=== Example: Single-DB ===") - example_shared_db(db_name="shared-traval-group11") + example_shared_db(db_name="shared-traval-group") + + print("\n=== Example: Single-DB-Complex ===") + example_complex_shared_db(db_name="shared-traval-group-complex-new") + + print("\n=== Example: Single-Community-DB-Complex ===") + example_complex_shared_db(db_name="paper", community=True) diff --git a/examples/data/config/tree_config_community.json b/examples/data/config/tree_config_community.json new file mode 100644 index 000000000..cbd2c8a07 --- /dev/null +++ b/examples/data/config/tree_config_community.json @@ -0,0 +1,50 @@ +{ + "extractor_llm": { + "backend": "ollama", + "config": { + "model_name_or_path": "qwen3:0.6b", + "temperature": 0.0, + "remove_think_prefix": true, + "max_tokens": 8192 + } + }, + "dispatcher_llm": { + "backend": "ollama", + "config": { + "model_name_or_path": "qwen3:0.6b", + "temperature": 0.0, + "remove_think_prefix": true, + "max_tokens": 8192 + } + }, + "embedder": { + "backend": "ollama", + "config": { + "model_name_or_path": "nomic-embed-text:latest" + } + }, + "graph_db": { + "backend": "neo4j-community", + "config": { + "uri": "bolt://localhost:7687", + "user": "neo4j", + "password": "12345678", + "db_name": "neo4j", + "user_name": "alice", + "use_multi_db": false, + "auto_create": false, + "embedding_dimension": 768, + "vec_config": { + "backend": "qdrant", + "config": { + "collection_name": "neo4j_vec_db", + "vector_dimension": 768, + "distance_metric": "cosine", + "host": "localhost", + "port": 6333 + } + } + } + }, + "reorganize": true +} diff --git a/examples/mem_os/simple_openapi_memos_neo4j_community.py b/examples/mem_os/simple_openapi_memos_neo4j_community.py new file mode 100644 index 000000000..aad1b8c77 --- /dev/null +++ b/examples/mem_os/simple_openapi_memos_neo4j_community.py @@ -0,0 +1,315 @@ +import os +import time +import uuid + +from datetime import datetime + +from dotenv import load_dotenv + +from memos.configs.mem_cube import GeneralMemCubeConfig +from memos.configs.mem_os import MOSConfig +from memos.mem_cube.general import GeneralMemCube +from memos.mem_os.main import MOS + + +load_dotenv() + +# 1. Create MOS Config and set openai config +print(f"πŸš€ [{datetime.now().strftime('%H:%M:%S')}] Starting to create MOS configuration...") +start_time = time.time() + +user_name = str(uuid.uuid4()) +print(user_name) + +# 1.1 Set openai config +openapi_config = { + "model_name_or_path": "gpt-4o-mini", + "temperature": 0.8, + "max_tokens": 1024, + "top_p": 0.9, + "top_k": 50, + "remove_think_prefix": True, + "api_key": os.getenv("OPENAI_API_KEY", "sk-xxxxx"), + "api_base": os.getenv("OPENAI_API_BASE", "https://api.openai.com/v1"), +} +embedder_config = { + "backend": "universal_api", + "config": { + "provider": "openai", + "api_key": os.getenv("OPENAI_API_KEY", "sk-xxxxx"), + "model_name_or_path": "text-embedding-3-large", + "base_url": os.getenv("OPENAI_API_BASE", "https://api.openai.com/v1"), + }, +} +EMBEDDING_DIMENSION = 3072 + +# 1.2 Set neo4j config +neo4j_uri = os.getenv("NEO4J_URI", "bolt://localhost:7687") + +# 1.3 Create MOS Config +config = { + "user_id": user_name, + "chat_model": { + "backend": "openai", + "config": openapi_config, + }, + "mem_reader": { + "backend": "simple_struct", + "config": { + "llm": { + "backend": "openai", + "config": openapi_config, + }, + "embedder": embedder_config, + "chunker": { + "backend": "sentence", + "config": { + "tokenizer_or_token_counter": "gpt2", + "chunk_size": 512, + "chunk_overlap": 128, + "min_sentences_per_chunk": 1, + }, + }, + }, + }, + "max_turns_window": 20, + "top_k": 5, + "enable_textual_memory": True, + "enable_activation_memory": False, + "enable_parametric_memory": False, +} + +mos_config = MOSConfig(**config) +# you can set PRO_MODE to True to enable CoT enhancement mos_config.PRO_MODE = True +mos = MOS(mos_config) + +print( + f"βœ… [{datetime.now().strftime('%H:%M:%S')}] MOS configuration created successfully, time elapsed: {time.time() - start_time:.2f}s\n" +) + +# 2. Initialize memory cube +print(f"πŸš€ [{datetime.now().strftime('%H:%M:%S')}] Starting to initialize MemCube configuration...") +start_time = time.time() + +config = GeneralMemCubeConfig.model_validate( + { + "user_id": user_name, + "cube_id": f"{user_name}", + "text_mem": { + "backend": "tree_text", + "config": { + "extractor_llm": { + "backend": "openai", + "config": openapi_config, + }, + "dispatcher_llm": { + "backend": "openai", + "config": openapi_config, + }, + "embedder": embedder_config, + "graph_db": { + "backend": "neo4j-community", + "config": { + "uri": neo4j_uri, + "user": "neo4j", + "password": "12345678", + "db_name": "neo4j", + "user_name": "alice", + "use_multi_db": False, + "auto_create": False, + "embedding_dimension": EMBEDDING_DIMENSION, + "vec_config": { + "backend": "qdrant", + "config": { + "collection_name": "neo4j_vec_db", + "vector_dimension": EMBEDDING_DIMENSION, + "distance_metric": "cosine", + "host": "localhost", + "port": 6333, + }, + }, + }, + }, + "reorganize": True, + }, + }, + "act_mem": {}, + "para_mem": {}, + }, +) + +print( + f"βœ… [{datetime.now().strftime('%H:%M:%S')}] MemCube configuration initialization completed, time elapsed: {time.time() - start_time:.2f}s\n" +) + +# 3. Initialize the MemCube with the configuration +print(f"πŸš€ [{datetime.now().strftime('%H:%M:%S')}] Starting to create MemCube instance...") +start_time = time.time() + +mem_cube = GeneralMemCube(config) +try: + mem_cube.dump(f"/tmp/{user_name}/") + print( + f"βœ… [{datetime.now().strftime('%H:%M:%S')}] MemCube created and saved successfully, time elapsed: {time.time() - start_time:.2f}s\n" + ) +except Exception as e: + print( + f"❌ [{datetime.now().strftime('%H:%M:%S')}] MemCube save failed: {e}, time elapsed: {time.time() - start_time:.2f}s\n" + ) + +# 4. Register the MemCube +print(f"πŸš€ [{datetime.now().strftime('%H:%M:%S')}] Starting to register MemCube...") +start_time = time.time() + +mos.register_mem_cube(f"/tmp/{user_name}", mem_cube_id=user_name) + +print( + f"βœ… [{datetime.now().strftime('%H:%M:%S')}] MemCube registration completed, time elapsed: {time.time() - start_time:.2f}s\n" +) + +# 5. Add, get, search memory +print(f"πŸš€ [{datetime.now().strftime('%H:%M:%S')}] Starting to add single memory...") +start_time = time.time() + +mos.add(memory_content="I like playing football.") + +print( + f"βœ… [{datetime.now().strftime('%H:%M:%S')}] Single memory added successfully, time elapsed: {time.time() - start_time:.2f}s" +) + +print(f"πŸš€ [{datetime.now().strftime('%H:%M:%S')}] Starting to get all memories...") +start_time = time.time() + +get_all_results = mos.get_all() + + +# Filter out embedding fields, keeping only necessary fields +def filter_memory_data(memories_data): + filtered_data = {} + for key, value in memories_data.items(): + if key == "text_mem": + filtered_data[key] = [] + for mem_group in value: + # Check if it's the new data structure (list of TextualMemoryItem objects) + if "memories" in mem_group and isinstance(mem_group["memories"], list): + # New data structure: directly a list of TextualMemoryItem objects + filtered_memories = [] + for memory_item in mem_group["memories"]: + # Create filtered dictionary + filtered_item = { + "id": memory_item.id, + "memory": memory_item.memory, + "metadata": {}, + } + # Filter metadata, excluding embedding + if hasattr(memory_item, "metadata") and memory_item.metadata: + for attr_name in dir(memory_item.metadata): + if not attr_name.startswith("_") and attr_name != "embedding": + attr_value = getattr(memory_item.metadata, attr_name) + if not callable(attr_value): + filtered_item["metadata"][attr_name] = attr_value + filtered_memories.append(filtered_item) + + filtered_group = { + "cube_id": mem_group.get("cube_id", ""), + "memories": filtered_memories, + } + filtered_data[key].append(filtered_group) + else: + # Old data structure: dictionary with nodes and edges + filtered_group = { + "memories": {"nodes": [], "edges": mem_group["memories"].get("edges", [])} + } + for node in mem_group["memories"].get("nodes", []): + filtered_node = { + "id": node.get("id"), + "memory": node.get("memory"), + "metadata": { + k: v + for k, v in node.get("metadata", {}).items() + if k != "embedding" + }, + } + filtered_group["memories"]["nodes"].append(filtered_node) + filtered_data[key].append(filtered_group) + else: + filtered_data[key] = value + return filtered_data + + +filtered_results = filter_memory_data(get_all_results) +print(f"Get all results after add memory: {filtered_results['text_mem'][0]['memories']}") + +print( + f"βœ… [{datetime.now().strftime('%H:%M:%S')}] Get all memories completed, time elapsed: {time.time() - start_time:.2f}s\n" +) + +# 6. Add messages +print(f"πŸš€ [{datetime.now().strftime('%H:%M:%S')}] Starting to add conversation messages...") +start_time = time.time() + +messages = [ + {"role": "user", "content": "I like playing football."}, + {"role": "assistant", "content": "yes football is my favorite game."}, +] +mos.add(messages) + +print( + f"βœ… [{datetime.now().strftime('%H:%M:%S')}] Conversation messages added successfully, time elapsed: {time.time() - start_time:.2f}s" +) + +print( + f"πŸš€ [{datetime.now().strftime('%H:%M:%S')}] Starting to get all memories (after adding messages)..." +) +start_time = time.time() + +get_all_results = mos.get_all() +filtered_results = filter_memory_data(get_all_results) +print(f"Get all results after add messages: {filtered_results}") + +print( + f"βœ… [{datetime.now().strftime('%H:%M:%S')}] Get all memories completed, time elapsed: {time.time() - start_time:.2f}s\n" +) + +# 7. Add document +print(f"πŸš€ [{datetime.now().strftime('%H:%M:%S')}] Starting to add document...") +start_time = time.time() +## 7.1 add pdf for ./tmp/data if use doc mem mos.add(doc_path="./tmp/data/") +start_time = time.time() + +get_all_results = mos.get_all() +filtered_results = filter_memory_data(get_all_results) +print(f"Get all results after add doc: {filtered_results}") + +print( + f"βœ… [{datetime.now().strftime('%H:%M:%S')}] Get all memories completed, time elapsed: {time.time() - start_time:.2f}s\n" +) + +# 8. Search +print(f"πŸš€ [{datetime.now().strftime('%H:%M:%S')}] Starting to search memories...") +start_time = time.time() + +search_results = mos.search(query="my favorite football game", user_id=user_name) +filtered_search_results = filter_memory_data(search_results) +print(f"Search results: {filtered_search_results}") + +print( + f"βœ… [{datetime.now().strftime('%H:%M:%S')}] Memory search completed, time elapsed: {time.time() - start_time:.2f}s\n" +) + +# 9. Chat +print(f"🎯 [{datetime.now().strftime('%H:%M:%S')}] Starting chat mode...") +while True: + user_input = input("πŸ‘€ [You] ").strip() + if user_input.lower() in ["quit", "exit"]: + break + + print() + chat_start_time = time.time() + response = mos.chat(user_input) + chat_duration = time.time() - chat_start_time + + print(f"πŸ€– [Assistant] {response}") + print(f"⏱️ [Response time: {chat_duration:.2f}s]\n") + +print("πŸ“’ [System] MemChat has stopped.") diff --git a/src/memos/configs/graph_db.py b/src/memos/configs/graph_db.py index 21baf4cd9..cca93fede 100644 --- a/src/memos/configs/graph_db.py +++ b/src/memos/configs/graph_db.py @@ -3,6 +3,7 @@ from pydantic import BaseModel, Field, field_validator, model_validator from memos.configs.base import BaseConfig +from memos.configs.vec_db import VectorDBConfigFactory class BaseGraphDBConfig(BaseConfig): @@ -79,12 +80,36 @@ def validate_config(self): return self +class Neo4jCommunityGraphDBConfig(Neo4jGraphDBConfig): + """ + Community edition config for Neo4j. + + Notes: + - Must set `use_multi_db = False` + - Must provide `user_name` for logical isolation + - Embedding vector DB config is required + """ + + vec_config: VectorDBConfigFactory = Field( + ..., description="Vector DB config for embedding search" + ) + + @model_validator(mode="after") + def validate_community(self): + if self.use_multi_db: + raise ValueError("Neo4j Community Edition does not support use_multi_db=True.") + if not self.user_name: + raise ValueError("Neo4j Community config requires user_name for logical isolation.") + return self + + class GraphDBConfigFactory(BaseModel): backend: str = Field(..., description="Backend for graph database") config: dict[str, Any] = Field(..., description="Configuration for the graph database backend") backend_to_class: ClassVar[dict[str, Any]] = { "neo4j": Neo4jGraphDBConfig, + "neo4j-community": Neo4jCommunityGraphDBConfig, } @field_validator("backend") diff --git a/src/memos/graph_dbs/factory.py b/src/memos/graph_dbs/factory.py index c100270da..c4365d16a 100644 --- a/src/memos/graph_dbs/factory.py +++ b/src/memos/graph_dbs/factory.py @@ -3,6 +3,7 @@ from memos.configs.graph_db import GraphDBConfigFactory from memos.graph_dbs.base import BaseGraphDB from memos.graph_dbs.neo4j import Neo4jGraphDB +from memos.graph_dbs.neo4j_community import Neo4jCommunityGraphDB class GraphStoreFactory(BaseGraphDB): @@ -10,6 +11,7 @@ class GraphStoreFactory(BaseGraphDB): backend_to_class: ClassVar[dict[str, Any]] = { "neo4j": Neo4jGraphDB, + "neo4j-community": Neo4jCommunityGraphDB, } @classmethod diff --git a/src/memos/graph_dbs/neo4j.py b/src/memos/graph_dbs/neo4j.py index 27ada91e5..1489ff541 100644 --- a/src/memos/graph_dbs/neo4j.py +++ b/src/memos/graph_dbs/neo4j.py @@ -12,18 +12,6 @@ logger = get_logger(__name__) -def _parse_node(node_data: dict[str, Any]) -> dict[str, Any]: - node = node_data.copy() - - # Convert Neo4j datetime to string - for time_field in ("created_at", "updated_at"): - if time_field in node and hasattr(node[time_field], "isoformat"): - node[time_field] = node[time_field].isoformat() - node.pop("user_name", None) - - return {"id": node.pop("id"), "memory": node.pop("memory", ""), "metadata": node} - - def _compose_node(item: dict[str, Any]) -> tuple[str, str, dict[str, Any]]: node_id = item["id"] memory = item["memory"] @@ -353,7 +341,7 @@ def get_node(self, id: str) -> dict[str, Any] | None: with self.driver.session(database=self.db_name) as session: record = session.run(query, params).single() - return _parse_node(dict(record["n"])) if record else None + return self._parse_node(dict(record["n"])) if record else None def get_nodes(self, ids: list[str]) -> list[dict[str, Any]]: """ @@ -381,7 +369,7 @@ def get_nodes(self, ids: list[str]) -> list[dict[str, Any]]: with self.driver.session(database=self.db_name) as session: results = session.run(query, params) - return [_parse_node(dict(record["n"])) for record in results] + return [self._parse_node(dict(record["n"])) for record in results] def get_edges(self, id: str, type: str = "ANY", direction: str = "ANY") -> list[dict[str, str]]: """ @@ -497,7 +485,7 @@ def get_neighbors_by_tag( with self.driver.session(database=self.db_name) as session: result = session.run(query, params) - return [_parse_node(dict(record["n"])) for record in result] + return [self._parse_node(dict(record["n"])) for record in result] def get_children_with_embeddings(self, id: str) -> list[dict[str, Any]]: where_user = "" @@ -579,8 +567,8 @@ def get_subgraph( if not centers or centers[0] is None: return {"core_node": None, "neighbors": [], "edges": []} - core_node = _parse_node(dict(centers[0])) - neighbors = [_parse_node(dict(n)) for n in record["neighbors"] if n] + core_node = self._parse_node(dict(centers[0])) + neighbors = [self._parse_node(dict(n)) for n in record["neighbors"] if n] edges = [] for rel_chain in record["rels"]: for rel in rel_chain: @@ -863,7 +851,7 @@ def export_graph(self) -> dict[str, Any]: params["user_name"] = self.config.user_name node_result = session.run(f"{node_query} RETURN n", params) - nodes = [_parse_node(dict(record["n"])) for record in node_result] + nodes = [self._parse_node(dict(record["n"])) for record in node_result] # Export edges edge_result = session.run( @@ -950,7 +938,7 @@ def get_all_memory_items(self, scope: str) -> list[dict]: with self.driver.session(database=self.db_name) as session: results = session.run(query, params) - return [_parse_node(dict(record["n"])) for record in results] + return [self._parse_node(dict(record["n"])) for record in results] def get_structure_optimization_candidates(self, scope: str) -> list[dict]: """ @@ -977,7 +965,9 @@ def get_structure_optimization_candidates(self, scope: str) -> list[dict]: with self.driver.session(database=self.db_name) as session: results = session.run(query, params) - return [_parse_node({"id": record["id"], **dict(record["node"])}) for record in results] + return [ + self._parse_node({"id": record["id"], **dict(record["node"])}) for record in results + ] def drop_database(self) -> None: """ @@ -1100,3 +1090,14 @@ def _index_exists(self, index_name: str) -> bool: if record["name"] == index_name: return True return False + + def _parse_node(self, node_data: dict[str, Any]) -> dict[str, Any]: + node = node_data.copy() + + # Convert Neo4j datetime to string + for time_field in ("created_at", "updated_at"): + if time_field in node and hasattr(node[time_field], "isoformat"): + node[time_field] = node[time_field].isoformat() + node.pop("user_name", None) + + return {"id": node.pop("id"), "memory": node.pop("memory", ""), "metadata": node} diff --git a/src/memos/graph_dbs/neo4j_community.py b/src/memos/graph_dbs/neo4j_community.py new file mode 100644 index 000000000..98d9723bb --- /dev/null +++ b/src/memos/graph_dbs/neo4j_community.py @@ -0,0 +1,300 @@ +from typing import Any + +from memos.configs.graph_db import Neo4jGraphDBConfig +from memos.graph_dbs.neo4j import Neo4jGraphDB, _prepare_node_metadata +from memos.log import get_logger +from memos.vec_dbs.factory import VecDBFactory +from memos.vec_dbs.item import VecDBItem + + +logger = get_logger(__name__) + + +class Neo4jCommunityGraphDB(Neo4jGraphDB): + """ + Neo4j Community Edition graph memory store. + + Note: + This class avoids Enterprise-only features: + - No multi-database support + - No vector index + - No CREATE DATABASE + """ + + def __init__(self, config: Neo4jGraphDBConfig): + assert config.auto_create is False + assert config.use_multi_db is False + # Init vector database + self.vec_db = VecDBFactory.from_config(config.vec_config) + # Call parent init + super().__init__(config) + + def create_index( + self, + label: str = "Memory", + vector_property: str = "embedding", + dimensions: int = 1536, + index_name: str = "memory_vector_index", + ) -> None: + """ + Create the vector index for embedding and datetime indexes for created_at and updated_at fields. + """ + # Create indexes + self._create_basic_property_indexes() + + def add_node(self, id: str, memory: str, metadata: dict[str, Any]) -> None: + if not self.config.use_multi_db and self.config.user_name: + metadata["user_name"] = self.config.user_name + + # Safely process metadata + metadata = _prepare_node_metadata(metadata) + + # Extract required fields + embedding = metadata.pop("embedding", None) + if embedding is None: + raise ValueError(f"Missing 'embedding' in metadata for node {id}") + + # Merge node and set metadata + created_at = metadata.pop("created_at") + updated_at = metadata.pop("updated_at") + vector_sync_status = "success" + + try: + # Write to Vector DB + item = VecDBItem( + id=id, + vector=embedding, + payload={ + "memory": memory, + "vector_sync": vector_sync_status, + **metadata, # unpack all metadata keys to top-level + }, + ) + self.vec_db.add([item]) + except Exception as e: + logger.warning(f"[VecDB] Vector insert failed for node {id}: {e}") + vector_sync_status = "failed" + + metadata["vector_sync"] = vector_sync_status + query = """ + MERGE (n:Memory {id: $id}) + SET n.memory = $memory, + n.created_at = datetime($created_at), + n.updated_at = datetime($updated_at), + n += $metadata + """ + with self.driver.session(database=self.db_name) as session: + session.run( + query, + id=id, + memory=memory, + created_at=created_at, + updated_at=updated_at, + metadata=metadata, + ) + + def get_children_with_embeddings(self, id: str) -> list[dict[str, Any]]: + where_user = "" + params = {"id": id} + + if not self.config.use_multi_db and self.config.user_name: + where_user = "AND p.user_name = $user_name AND c.user_name = $user_name" + params["user_name"] = self.config.user_name + + query = f""" + MATCH (p:Memory)-[:PARENT]->(c:Memory) + WHERE p.id = $id {where_user} + RETURN c.id AS id, c.memory AS memory + """ + + with self.driver.session(database=self.db_name) as session: + result = session.run(query, params) + child_nodes = [{"id": r["id"], "memory": r["memory"]} for r in result] + + # Get embeddings from vector DB + ids = [n["id"] for n in child_nodes] + vec_items = {v.id: v.vector for v in self.vec_db.get_by_ids(ids)} + + # Merge results + for node in child_nodes: + node["embedding"] = vec_items.get(node["id"]) + + return child_nodes + + # Search / recall operations + def search_by_embedding( + self, + vector: list[float], + top_k: int = 5, + scope: str | None = None, + status: str | None = None, + threshold: float | None = None, + ) -> list[dict]: + """ + Retrieve node IDs based on vector similarity using external vector DB. + + Args: + vector (list[float]): The embedding vector representing query semantics. + top_k (int): Number of top similar nodes to retrieve. + scope (str, optional): Memory type filter (e.g., 'WorkingMemory', 'LongTermMemory'). + status (str, optional): Node status filter (e.g., 'activated', 'archived'). + threshold (float, optional): Minimum similarity score threshold (0 ~ 1). + + Returns: + list[dict]: A list of dicts with 'id' and 'score', ordered by similarity. + + Notes: + - This method uses an external vector database (not Neo4j) to perform the search. + - If 'scope' is provided, it restricts results to nodes with matching memory_type. + - If 'status' is provided, it further filters nodes by status. + - If 'threshold' is provided, only results with score >= threshold will be returned. + - The returned IDs can be used to fetch full node data from Neo4j if needed. + """ + # Build VecDB filter + vec_filter = {} + if scope: + vec_filter["memory_type"] = scope + if status: + vec_filter["status"] = status + vec_filter["vector_sync"] = "success" + vec_filter["user_name"] = self.config.user_name + + # Perform vector search + results = self.vec_db.search(query_vector=vector, top_k=top_k, filter=vec_filter) + + # Filter by threshold + if threshold is not None: + results = [r for r in results if r.score is None or r.score >= threshold] + + # Return consistent format + return [{"id": r.id, "score": r.score} for r in results] + + def get_all_memory_items(self, scope: str) -> list[dict]: + """ + Retrieve all memory items of a specific memory_type. + + Args: + scope (str): Must be one of 'WorkingMemory', 'LongTermMemory', or 'UserMemory'. + + Returns: + list[dict]: Full list of memory items under this scope. + """ + if scope not in {"WorkingMemory", "LongTermMemory", "UserMemory"}: + raise ValueError(f"Unsupported memory type scope: {scope}") + + where_clause = "WHERE n.memory_type = $scope" + params = {"scope": scope} + + if not self.config.use_multi_db and self.config.user_name: + where_clause += " AND n.user_name = $user_name" + params["user_name"] = self.config.user_name + + query = f""" + MATCH (n:Memory) + {where_clause} + RETURN n + """ + + with self.driver.session(database=self.db_name) as session: + results = session.run(query, params) + return [self._parse_node(dict(record["n"])) for record in results] + + def clear(self) -> None: + """ + Clear the entire graph if the target database exists. + """ + # Step 1: clear Neo4j part via parent logic + super().clear() + + # Step2: Clear the vector db + try: + items = self.vec_db.get_by_filter({"user_name": self.config.user_name}) + if items: + self.vec_db.delete([item.id for item in items]) + logger.info(f"Cleared {len(items)} vectors for user '{self.config.user_name}'.") + else: + logger.info(f"No vectors to clear for user '{self.config.user_name}'.") + except Exception as e: + logger.warning(f"Failed to clear vector DB for user '{self.config.user_name}': {e}") + + def drop_database(self) -> None: + """ + Permanently delete the entire database this instance is using. + WARNING: This operation is destructive and cannot be undone. + """ + raise ValueError( + f"Refusing to drop protected database: {self.db_name} in " + f"Shared Database Multi-Tenant mode" + ) + + # Avoid enterprise feature + def _ensure_database_exists(self): + pass + + def _create_basic_property_indexes(self) -> None: + """ + Create standard B-tree indexes on memory_type, created_at, + and updated_at fields. + Create standard B-tree indexes on user_name when use Shared Database + Multi-Tenant Mode + """ + # Step 1: Neo4j indexes + try: + with self.driver.session(database=self.db_name) as session: + session.run(""" + CREATE INDEX memory_type_index IF NOT EXISTS + FOR (n:Memory) ON (n.memory_type) + """) + logger.debug("Index 'memory_type_index' ensured.") + + session.run(""" + CREATE INDEX memory_created_at_index IF NOT EXISTS + FOR (n:Memory) ON (n.created_at) + """) + logger.debug("Index 'memory_created_at_index' ensured.") + + session.run(""" + CREATE INDEX memory_updated_at_index IF NOT EXISTS + FOR (n:Memory) ON (n.updated_at) + """) + logger.debug("Index 'memory_updated_at_index' ensured.") + + if not self.config.use_multi_db and self.config.user_name: + session.run( + """ + CREATE INDEX memory_user_name_index IF NOT EXISTS + FOR (n:Memory) ON (n.user_name) + """ + ) + logger.debug("Index 'memory_user_name_index' ensured.") + except Exception as e: + logger.warning(f"Failed to create basic property indexes: {e}") + + # Step 2: VectorDB indexes + try: + if hasattr(self.vec_db, "ensure_payload_indexes"): + self.vec_db.ensure_payload_indexes(["user_name", "memory_type", "status"]) + else: + logger.debug("VecDB does not support payload index creation; skipping.") + except Exception as e: + logger.warning(f"Failed to create VecDB payload indexes: {e}") + + def _parse_node(self, node_data: dict[str, Any]) -> dict[str, Any]: + """Parse Neo4j node and optionally fetch embedding from vector DB.""" + node = node_data.copy() + + # Convert Neo4j datetime to string + for time_field in ("created_at", "updated_at"): + if time_field in node and hasattr(node[time_field], "isoformat"): + node[time_field] = node[time_field].isoformat() + node.pop("user_name", None) + + new_node = {"id": node.pop("id"), "memory": node.pop("memory", ""), "metadata": node} + try: + vec_item = self.vec_db.get_by_id(new_node["id"]) + if vec_item and vec_item.vector: + new_node["metadata"]["embedding"] = vec_item.vector + except Exception as e: + logger.warning(f"Failed to fetch vector for node {new_node['id']}: {e}") + new_node["metadata"]["embedding"] = None + return new_node diff --git a/src/memos/memories/textual/tree_text_memory/retrieve/recall.py b/src/memos/memories/textual/tree_text_memory/retrieve/recall.py index 8a3fc4b5c..36a0b5fee 100644 --- a/src/memos/memories/textual/tree_text_memory/retrieve/recall.py +++ b/src/memos/memories/textual/tree_text_memory/retrieve/recall.py @@ -56,7 +56,6 @@ def retrieve( # Step 3: Merge and deduplicate results combined = {item.id: item for item in graph_results + vector_results} - # Debug: ζ‰“ε°εœ¨ graph_results δΈ­δ½†δΈεœ¨ combined δΈ­ηš„ id graph_ids = {item.id for item in graph_results} combined_ids = set(combined.keys()) lost_ids = graph_ids - combined_ids diff --git a/src/memos/vec_dbs/base.py b/src/memos/vec_dbs/base.py index 2ffa3b539..ee1bfb3ca 100644 --- a/src/memos/vec_dbs/base.py +++ b/src/memos/vec_dbs/base.py @@ -55,6 +55,10 @@ def search( def get_by_id(self, id: str) -> VecDBItem | None: """Get an item from the vector database.""" + @abstractmethod + def get_by_ids(self, ids: list[str]) -> list[VecDBItem]: + """Get multiple items by their IDs.""" + @abstractmethod def get_by_filter(self, filter: dict[str, Any]) -> list[VecDBItem]: """ @@ -103,3 +107,11 @@ def upsert(self, data: list[VecDBItem | dict[str, Any]]) -> None: @abstractmethod def delete(self, ids: list[str]) -> None: """Delete items from the vector database.""" + + @abstractmethod + def ensure_payload_indexes(self, fields: list[str]) -> None: + """ + Create payload indexes for specified fields in the collection. + Args: + fields (list[str]): List of field names to index (as keyword). + """ diff --git a/src/memos/vec_dbs/qdrant.py b/src/memos/vec_dbs/qdrant.py index c37f113bd..a0ebf1d80 100644 --- a/src/memos/vec_dbs/qdrant.py +++ b/src/memos/vec_dbs/qdrant.py @@ -278,6 +278,25 @@ def update(self, id: str, data: VecDBItem | dict[str, Any]) -> None: collection_name=self.config.collection_name, payload=data.payload, points=[id] ) + def ensure_payload_indexes(self, fields: list[str]) -> None: + """ + Create payload indexes for specified fields in the collection. + This is idempotent: it will skip if index already exists. + + Args: + fields (list[str]): List of field names to index (as keyword). + """ + for field in fields: + try: + self.client.create_payload_index( + collection_name=self.config.collection_name, + field_name=field, + field_schema="keyword", # Could be extended in future + ) + logger.debug(f"Qdrant payload index on '{field}' ensured.") + except Exception as e: + logger.warning(f"Failed to create payload index on '{field}': {e}") + def upsert(self, data: list[VecDBItem | dict[str, Any]]) -> None: """ Add or update data in the vector database.