diff --git a/.gitignore b/.gitignore index 184181828..792869e14 100644 --- a/.gitignore +++ b/.gitignore @@ -6,6 +6,8 @@ tmp/ **/tmp_data/ # evaluation data +*.csv +*.jsonl evaluation/*tmp/ evaluation/results evaluation/.env @@ -13,6 +15,7 @@ evaluation/.env evaluation/configs/* **tree_textual_memory_locomo** .env +evaluation/scripts/personamem # Byte-compiled / optimized / DLL files __pycache__/ diff --git a/docker/Dockerfile b/docker/Dockerfile index ea158b111..29636881c 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -1,5 +1,7 @@ +# Base image FROM python:3.11-slim +# Install dependencies RUN apt-get update && apt-get install -y \ gcc \ g++ \ @@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y \ curl \ && rm -rf /var/lib/apt/lists/* +# Set working directory WORKDIR /app +# Set Hugging Face mirror ENV HF_ENDPOINT=https://hf-mirror.com -COPY requirements.txt . +# Install Python packages +COPY docker/requirements.txt . RUN pip install --upgrade pip && pip install --no-cache-dir -r requirements.txt -RUN pip install chonkie +# Copy application code +COPY docker/ ./docker/ +COPY src/ ./src/ -COPY ../. . +# Set Python import path ENV PYTHONPATH=/app/src +# Expose port EXPOSE 8000 + +# Start the docker CMD ["uvicorn", "memos.api.product_api:app", "--host", "0.0.0.0", "--port", "8000", "--reload"] diff --git a/docker/docker-compose.yml b/docker/docker-compose.yml index 64d383a25..d8998b6f7 100644 --- a/docker/docker-compose.yml +++ b/docker/docker-compose.yml @@ -2,10 +2,10 @@ name: memos-dev services: memos: - container_name: memos-api-server + container_name: memos-api-docker build: context: .. - dockerfile: Dockerfile + dockerfile: docker/Dockerfile ports: - "8000:8000" env_file: @@ -15,14 +15,16 @@ services: - qdrant environment: - PYTHONPATH=/app/src + - HF_ENDPOINT=https://hf-mirror.com volumes: - - .:/app + - ../src:/app/src + - .:/app/docker networks: - memos_network neo4j: image: neo4j:5.26.4 - container_name: neo4j-server + container_name: neo4j-docker ports: - "7474:7474" # HTTP - "7687:7687" # Bolt @@ -43,7 +45,7 @@ services: qdrant: image: qdrant/qdrant:v1.15.0 - container_name: qdrant-server + container_name: qdrant-docker ports: - "6333:6333" # REST API - "6334:6334" # gRPC API diff --git a/docker/requirements.txt b/docker/requirements.txt index 596e94b47..211ec3cae 100644 --- a/docker/requirements.txt +++ b/docker/requirements.txt @@ -2,24 +2,33 @@ annotated-types==0.7.0 ; python_version >= "3.10" and python_version < "4.0" anyio==4.9.0 ; python_version >= "3.10" and python_version < "4.0" attrs==25.3.0 ; python_version >= "3.10" and python_version < "4.0" authlib==1.6.0 ; python_version >= "3.10" and python_version < "4.0" +beautifulsoup4==4.13.4 ; python_version >= "3.10" and python_version < "4.0" certifi==2025.7.14 ; python_version >= "3.10" and python_version < "4.0" cffi==1.17.1 ; python_version >= "3.10" and python_version < "4.0" and platform_python_implementation != "PyPy" +cfgv==3.4.0 ; python_version >= "3.10" and python_version < "4.0" charset-normalizer==3.4.2 ; python_version >= "3.10" and python_version < "4.0" +chonkie==1.1.1 ; python_version >= "3.10" and python_version < "4.0" click==8.2.1 ; python_version >= "3.10" and python_version < "4.0" +cobble==0.1.4 ; python_version >= "3.10" and python_version < "4.0" colorama==0.4.6 ; python_version >= "3.10" and python_version < "4.0" and (platform_system == "Windows" or sys_platform == "win32") +coloredlogs==15.0.1 ; python_version >= "3.10" and python_version < "4.0" cryptography==45.0.5 ; python_version >= "3.10" and python_version < "4.0" cyclopts==3.22.2 ; python_version >= "3.10" and python_version < "4.0" +defusedxml==0.7.1 ; python_version >= "3.10" and python_version < "4.0" +distlib==0.4.0 ; python_version >= "3.10" and python_version < "4.0" distro==1.9.0 ; python_version >= "3.10" and python_version < "4.0" dnspython==2.7.0 ; python_version >= "3.10" and python_version < "4.0" docstring-parser==0.16 ; python_version >= "3.10" and python_version < "4.0" docutils==0.21.2 ; python_version >= "3.10" and python_version < "4.0" email-validator==2.2.0 ; python_version >= "3.10" and python_version < "4.0" +et-xmlfile==2.0.0 ; python_version >= "3.10" and python_version < "4.0" exceptiongroup==1.3.0 ; python_version >= "3.10" and python_version < "4.0" fastapi-cli==0.0.8 ; python_version >= "3.10" and python_version < "4.0" fastapi-cloud-cli==0.1.4 ; python_version >= "3.10" and python_version < "4.0" fastapi==0.115.14 ; python_version >= "3.10" and python_version < "4.0" fastmcp==2.10.5 ; python_version >= "3.10" and python_version < "4.0" filelock==3.18.0 ; python_version >= "3.10" and python_version < "4.0" +flatbuffers==25.2.10 ; python_version >= "3.10" and python_version < "4.0" fsspec==2025.7.0 ; python_version >= "3.10" and python_version < "4.0" greenlet==3.2.3 ; python_version >= "3.10" and python_version < "3.14" and (platform_machine == "aarch64" or platform_machine == "ppc64le" or platform_machine == "x86_64" or platform_machine == "amd64" or platform_machine == "AMD64" or platform_machine == "win32" or platform_machine == "WIN32") h11==0.16.0 ; python_version >= "3.10" and python_version < "4.0" @@ -29,24 +38,44 @@ httptools==0.6.4 ; python_version >= "3.10" and python_version < "4.0" httpx-sse==0.4.1 ; python_version >= "3.10" and python_version < "4.0" httpx==0.28.1 ; python_version >= "3.10" and python_version < "4.0" huggingface-hub==0.33.4 ; python_version >= "3.10" and python_version < "4.0" +humanfriendly==10.0 ; python_version >= "3.10" and python_version < "4.0" +identify==2.6.12 ; python_version >= "3.10" and python_version < "4.0" idna==3.10 ; python_version >= "3.10" and python_version < "4.0" +iniconfig==2.1.0 ; python_version >= "3.10" and python_version < "4.0" itsdangerous==2.2.0 ; python_version >= "3.10" and python_version < "4.0" jinja2==3.1.6 ; python_version >= "3.10" and python_version < "4.0" jiter==0.10.0 ; python_version >= "3.10" and python_version < "4.0" joblib==1.5.1 ; python_version >= "3.10" and python_version < "4.0" jsonschema-specifications==2025.4.1 ; python_version >= "3.10" and python_version < "4.0" jsonschema==4.24.1 ; python_version >= "3.10" and python_version < "4.0" +lxml==6.0.0 ; python_version >= "3.10" and python_version < "4.0" +magika==0.6.2 ; python_version >= "3.10" and python_version < "4.0" +mammoth==1.9.1 ; python_version >= "3.10" and python_version < "4.0" markdown-it-py==3.0.0 ; python_version >= "3.10" and python_version < "4.0" +markdownify==1.1.0 ; python_version >= "3.10" and python_version < "4.0" +markitdown==0.1.2 ; python_version >= "3.10" and python_version < "4.0" markupsafe==3.0.2 ; python_version >= "3.10" and python_version < "4.0" mcp==1.12.0 ; python_version >= "3.10" and python_version < "4.0" mdurl==0.1.2 ; python_version >= "3.10" and python_version < "4.0" +mpmath==1.3.0 ; python_version >= "3.10" and python_version < "4.0" +neo4j==5.28.1 ; python_version >= "3.10" and python_version < "4.0" +nodeenv==1.9.1 ; python_version >= "3.10" and python_version < "4.0" numpy==2.2.6 ; python_version == "3.10" numpy==2.3.1 ; python_version >= "3.11" and python_version < "4.0" ollama==0.4.9 ; python_version >= "3.10" and python_version < "4.0" +onnxruntime==1.22.1 ; python_version >= "3.10" and python_version < "4.0" openai==1.97.0 ; python_version >= "3.10" and python_version < "4.0" openapi-pydantic==0.5.1 ; python_version >= "3.10" and python_version < "4.0" +openpyxl==3.1.5 ; python_version >= "3.10" and python_version < "4.0" orjson==3.11.0 ; python_version >= "3.10" and python_version < "4.0" packaging==25.0 ; python_version >= "3.10" and python_version < "4.0" +pandas==2.3.1 ; python_version >= "3.10" and python_version < "4.0" +pdfminer-six==20250506 ; python_version >= "3.10" and python_version < "4.0" +pillow==11.3.0 ; python_version >= "3.10" and python_version < "4.0" +platformdirs==4.3.8 ; python_version >= "3.10" and python_version < "4.0" +pluggy==1.6.0 ; python_version >= "3.10" and python_version < "4.0" +pre-commit==4.2.0 ; python_version >= "3.10" and python_version < "4.0" +protobuf==6.31.1 ; python_version >= "3.10" and python_version < "4.0" pycparser==2.22 ; python_version >= "3.10" and python_version < "4.0" and platform_python_implementation != "PyPy" pydantic-core==2.33.2 ; python_version >= "3.10" and python_version < "4.0" pydantic-extra-types==2.10.5 ; python_version >= "3.10" and python_version < "4.0" @@ -54,8 +83,14 @@ pydantic-settings==2.10.1 ; python_version >= "3.10" and python_version < "4.0" pydantic==2.11.7 ; python_version >= "3.10" and python_version < "4.0" pygments==2.19.2 ; python_version >= "3.10" and python_version < "4.0" pyperclip==1.9.0 ; python_version >= "3.10" and python_version < "4.0" +pyreadline3==3.5.4 ; python_version >= "3.10" and python_version < "4.0" and sys_platform == "win32" +pytest-asyncio==0.23.8 ; python_version >= "3.10" and python_version < "4.0" +pytest==8.4.1 ; python_version >= "3.10" and python_version < "4.0" +python-dateutil==2.9.0.post0 ; python_version >= "3.10" and python_version < "4.0" python-dotenv==1.1.1 ; python_version >= "3.10" and python_version < "4.0" python-multipart==0.0.20 ; python_version >= "3.10" and python_version < "4.0" +python-pptx==1.0.2 ; python_version >= "3.10" and python_version < "4.0" +pytz==2025.2 ; python_version >= "3.10" and python_version < "4.0" pywin32==311 ; python_version >= "3.10" and python_version < "4.0" and (platform_system == "Windows" or sys_platform == "win32") pyyaml==6.0.2 ; python_version >= "3.10" and python_version < "4.0" referencing==0.36.2 ; python_version >= "3.10" and python_version < "4.0" @@ -66,27 +101,37 @@ rich-toolkit==0.14.8 ; python_version >= "3.10" and python_version < "4.0" rich==14.0.0 ; python_version >= "3.10" and python_version < "4.0" rignore==0.6.2 ; python_version >= "3.10" and python_version < "4.0" rpds-py==0.26.0 ; python_version >= "3.10" and python_version < "4.0" +ruff==0.11.13 ; python_version >= "3.10" and python_version < "4.0" safetensors==0.5.3 ; python_version >= "3.10" and python_version < "4.0" +schedule==1.2.2 ; python_version >= "3.10" and python_version < "4.0" scikit-learn==1.7.0 ; python_version >= "3.10" and python_version < "4.0" scipy==1.15.3 ; python_version == "3.10" scipy==1.16.0 ; python_version >= "3.11" and python_version < "4.0" sentry-sdk==2.33.0 ; python_version >= "3.10" and python_version < "4.0" shellingham==1.5.4 ; python_version >= "3.10" and python_version < "4.0" +six==1.17.0 ; python_version >= "3.10" and python_version < "4.0" sniffio==1.3.1 ; python_version >= "3.10" and python_version < "4.0" +soupsieve==2.7 ; python_version >= "3.10" and python_version < "4.0" sqlalchemy==2.0.41 ; python_version >= "3.10" and python_version < "4.0" sse-starlette==2.4.1 ; python_version >= "3.10" and python_version < "4.0" starlette==0.46.2 ; python_version >= "3.10" and python_version < "4.0" +sympy==1.14.0 ; python_version >= "3.10" and python_version < "4.0" tenacity==9.1.2 ; python_version >= "3.10" and python_version < "4.0" threadpoolctl==3.6.0 ; python_version >= "3.10" and python_version < "4.0" tokenizers==0.21.2 ; python_version >= "3.10" and python_version < "4.0" +tomli==2.2.1 ; python_version == "3.10" tqdm==4.67.1 ; python_version >= "3.10" and python_version < "4.0" transformers==4.53.2 ; python_version >= "3.10" and python_version < "4.0" typer==0.16.0 ; python_version >= "3.10" and python_version < "4.0" typing-extensions==4.14.1 ; python_version >= "3.10" and python_version < "4.0" typing-inspection==0.4.1 ; python_version >= "3.10" and python_version < "4.0" +tzdata==2025.2 ; python_version >= "3.10" and python_version < "4.0" ujson==5.10.0 ; python_version >= "3.10" and python_version < "4.0" urllib3==2.5.0 ; python_version >= "3.10" and python_version < "4.0" uvicorn==0.35.0 ; python_version >= "3.10" and python_version < "4.0" uvloop==0.21.0 ; python_version >= "3.10" and python_version < "4.0" and platform_python_implementation != "PyPy" and sys_platform != "win32" and sys_platform != "cygwin" +virtualenv==20.31.2 ; python_version >= "3.10" and python_version < "4.0" watchfiles==1.1.0 ; python_version >= "3.10" and python_version < "4.0" websockets==15.0.1 ; python_version >= "3.10" and python_version < "4.0" +xlrd==2.0.2 ; python_version >= "3.10" and python_version < "4.0" +xlsxwriter==3.2.5 ; python_version >= "3.10" and python_version < "4.0" diff --git a/docs/openapi.json b/docs/openapi.json index 15e834d04..b4193dd67 100644 --- a/docs/openapi.json +++ b/docs/openapi.json @@ -884,7 +884,7 @@ "type": "string", "title": "Session Id", "description": "Session ID for the MOS. This is used to distinguish between different dialogue", - "default": "a47d75a0-5ee8-473f-86c4-3f09073fd59f" + "default": "842877f4-c3f7-4c22-ad38-5950026870fe" }, "chat_model": { "$ref": "#/components/schemas/LLMConfigFactory", @@ -905,6 +905,10 @@ ], "description": "Memory scheduler configuration for managing memory operations" }, + "user_manager": { + "$ref": "#/components/schemas/UserManagerConfigFactory", + "description": "User manager configuration for database operations" + }, "max_turns_window": { "type": "integer", "title": "Max Turns Window", @@ -1370,6 +1374,25 @@ "title": "UserListResponse", "description": "Response model for user list operations." }, + "UserManagerConfigFactory": { + "properties": { + "backend": { + "type": "string", + "title": "Backend", + "description": "Backend for user manager", + "default": "sqlite" + }, + "config": { + "additionalProperties": true, + "type": "object", + "title": "Config", + "description": "Configuration for the user manager backend" + } + }, + "type": "object", + "title": "UserManagerConfigFactory", + "description": "Factory for user manager configurations." + }, "UserResponse": { "properties": { "code": { diff --git a/evaluation/data/personamem/.gitkeep b/evaluation/data/personamem/.gitkeep new file mode 100644 index 000000000..e69de29bb diff --git a/evaluation/scripts/locomo/locomo_eval.py b/evaluation/scripts/locomo/locomo_eval.py index c6adbd61c..25d2a847e 100644 --- a/evaluation/scripts/locomo/locomo_eval.py +++ b/evaluation/scripts/locomo/locomo_eval.py @@ -32,7 +32,6 @@ except Exception as e: print(f"Warning: Failed to download NLTK resources: {e}") - try: sentence_model_name = "Qwen/Qwen3-Embedding-0.6B" sentence_model = SentenceTransformer(sentence_model_name) @@ -363,8 +362,7 @@ async def limited_task(task): parser.add_argument( "--lib", type=str, - choices=["zep", "memos", "mem0", "mem0_graph", "langmem", "openai"], - help="Specify the memory framework (zep or memos or mem0 or mem0_graph)", + choices=["zep", "memos", "mem0", "mem0_graph", "openai", "memos-api", "memobase"], ) parser.add_argument( "--version", diff --git a/evaluation/scripts/locomo/locomo_ingestion.py b/evaluation/scripts/locomo/locomo_ingestion.py index f3837002d..ae5e57c87 100644 --- a/evaluation/scripts/locomo/locomo_ingestion.py +++ b/evaluation/scripts/locomo/locomo_ingestion.py @@ -1,7 +1,26 @@ +import os +import sys +import uuid + + +sys.path.insert( + 0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))) +) +sys.path.insert( + 0, + os.path.join( + os.path.dirname( + os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + ), + "evaluation", + "scripts", + ), +) + import argparse import concurrent.futures import json -import os +import threading import time from datetime import datetime, timezone @@ -10,7 +29,9 @@ from dotenv import load_dotenv from mem0 import MemoryClient +from memobase import ChatBlob from tqdm import tqdm +from utils.client import memobase_client, memos_client from zep_cloud.client import Zep from memos.configs.mem_cube import GeneralMemCubeConfig @@ -93,7 +114,34 @@ def get_client(frame: str, user_id: str | None = None, version: str = "default") return mos -def ingest_session(client, session, frame, metadata, revised_client=None): +def string_to_uuid(s: str, salt="memobase_client") -> str: + return str(uuid.uuid5(uuid.NAMESPACE_DNS, s + salt)) + + +def memobase_add_memory(user, message, retries=3): + for attempt in range(retries): + try: + _ = user.insert(ChatBlob(messages=message), sync=True) + return + except Exception as e: + if attempt < retries - 1: + time.sleep(1) + continue + else: + raise e + + +def memobase_add_memories_for_speaker(client, speaker, messages): + real_uid = string_to_uuid(speaker) + u = client.get_or_create_user(real_uid) + for i in range(0, len(messages), 2): + batch_messages = messages[i : i + 2] + memobase_add_memory(u, batch_messages) + print(f"[{i + 1}/{len(messages)}] Added messages for {speaker} successfully.") + u.flush(sync=True) + + +def ingest_session(client, session, frame, version, metadata, revised_client=None): session_date = metadata["session_date"] date_format = "%I:%M %p on %d %B, %Y UTC" date_string = datetime.strptime(session_date, date_format).replace(tzinfo=timezone.utc) @@ -125,7 +173,7 @@ def ingest_session(client, session, frame, metadata, revised_client=None): group_id=conv_id, ) - elif frame == "memos": + elif frame == "memos" or frame == "memos-api": messages = [] messages_reverse = [] @@ -149,16 +197,22 @@ def ingest_session(client, session, frame, metadata, revised_client=None): speaker_a_user_id = conv_id + "_speaker_a" speaker_b_user_id = conv_id + "_speaker_b" + if frame == "memos-api": + client.add(messages=messages, user_id=f"{speaker_a_user_id.replace('_', '')}{version}") - client.add( - messages=messages, - user_id=speaker_a_user_id, - ) + revised_client.add( + messages=messages_reverse, user_id=f"{speaker_b_user_id.replace('_', '')}{version}" + ) + elif frame == "memos": + client.add( + messages=messages, + user_id=speaker_a_user_id, + ) - revised_client.add( - messages=messages_reverse, - user_id=speaker_b_user_id, - ) + revised_client.add( + messages=messages_reverse, + user_id=speaker_b_user_id, + ) print(f"Added messages for {speaker_a_user_id} and {speaker_b_user_id} successfully.") elif frame == "mem0" or frame == "mem0_graph": @@ -217,6 +271,77 @@ def ingest_session(client, session, frame, metadata, revised_client=None): version="v2", enable_graph=True, ) + elif frame == "memobase": + print(f"Processing abc for {metadata['session_key']}") + messages = [] + messages_reverse = [] + + for chat in tqdm(session, desc=f"{metadata['session_key']}"): + data = chat.get("speaker") + ": " + chat.get("text") + + if chat.get("speaker") == metadata["speaker_a"]: + messages.append( + { + "role": "user", + "content": chat.get("text"), + "alias": metadata["speaker_a"], + "created_at": iso_date, + } + ) + messages_reverse.append( + { + "role": "assistant", + "content": chat.get("text"), + "alias": metadata["speaker_b"], + "created_at": iso_date, + } + ) + elif chat.get("speaker") == metadata["speaker_b"]: + messages.append( + { + "role": "assistant", + "content": chat.get("text"), + "alias": metadata["speaker_b"], + "created_at": iso_date, + } + ) + messages_reverse.append( + { + "role": "user", + "content": chat.get("text"), + "alias": metadata["speaker_a"], + "created_at": iso_date, + } + ) + else: + raise ValueError( + f"Unknown speaker {chat.get('speaker')} in session {metadata['session_key']}" + ) + + print({"context": data, "conv_id": conv_id, "created_at": iso_date}) + + thread_a = threading.Thread( + target=memobase_add_memories_for_speaker, + args=( + client, + metadata["speaker_a_user_id"], + messages, + ), + ) + + thread_b = threading.Thread( + target=memobase_add_memories_for_speaker, + args=( + client, + metadata["speaker_b_user_id"], + messages_reverse, + ), + ) + + thread_a.start() + thread_b.start() + thread_a.join() + thread_b.join() end_time = time.time() elapsed_time = round(end_time - start_time, 2) @@ -246,7 +371,19 @@ def process_user(conv_idx, frame, locomo_df, version, num_workers=1): speaker_b_user_id = conv_id + "_speaker_b" client = get_client("memos", speaker_a_user_id, version) revised_client = get_client("memos", speaker_b_user_id, version) - + elif frame == "memos-api": + conv_id = "locomo_exp_user_" + str(conv_idx) + speaker_a_user_id = conv_id + "_speaker_a" + speaker_b_user_id = conv_id + "_speaker_b" + client = memos_client(mode="api") + revised_client = memos_client(mode="api") + elif frame == "memobase": + client = memobase_client() + conv_id = "locomo_exp_user_" + str(conv_idx) + speaker_a_user_id = conv_id + "_speaker_a" + speaker_b_user_id = conv_id + "_speaker_b" + client.delete_user(string_to_uuid(speaker_a_user_id)) + client.delete_user(string_to_uuid(speaker_b_user_id)) sessions_to_process = [] for session_idx in range(max_session_count): session_key = f"session_{session_idx}" @@ -272,7 +409,7 @@ def process_user(conv_idx, frame, locomo_df, version, num_workers=1): with concurrent.futures.ThreadPoolExecutor(max_workers=num_workers) as executor: futures = { executor.submit( - ingest_session, client, session, frame, metadata, revised_client + ingest_session, client, session, frame, version, metadata, revised_client ): metadata["session_key"] for session, metadata in sessions_to_process } @@ -340,8 +477,7 @@ def main(frame, version="default", num_workers=4): parser.add_argument( "--lib", type=str, - choices=["zep", "memos", "mem0", "mem0_graph"], - help="Specify the memory framework (zep or memos or mem0 or mem0_graph)", + choices=["zep", "memos", "mem0", "mem0_graph", "memos-api", "memobase"], ) parser.add_argument( "--version", diff --git a/evaluation/scripts/locomo/locomo_metric.py b/evaluation/scripts/locomo/locomo_metric.py index 9335ec5ba..8ee18faaf 100644 --- a/evaluation/scripts/locomo/locomo_metric.py +++ b/evaluation/scripts/locomo/locomo_metric.py @@ -9,8 +9,7 @@ parser.add_argument( "--lib", type=str, - choices=["zep", "memos", "mem0", "mem0_graph", "langmem", "openai"], - help="Specify the memory framework (zep or memos or mem0 or mem0_graph)", + choices=["zep", "memos", "mem0", "mem0_graph", "openai", "memos-api", "memobase"], ) parser.add_argument( "--version", diff --git a/evaluation/scripts/locomo/locomo_responses.py b/evaluation/scripts/locomo/locomo_responses.py index 5d0374c2b..056b17163 100644 --- a/evaluation/scripts/locomo/locomo_responses.py +++ b/evaluation/scripts/locomo/locomo_responses.py @@ -124,8 +124,7 @@ async def main(frame, version="default"): parser.add_argument( "--lib", type=str, - choices=["zep", "memos", "mem0", "mem0_graph", "openai"], - help="Specify the memory framework (zep or memos or mem0 or mem0_graph)", + choices=["zep", "memos", "mem0", "mem0_graph", "openai", "memos-api", "memobase"], ) parser.add_argument( "--version", diff --git a/evaluation/scripts/locomo/locomo_search.py b/evaluation/scripts/locomo/locomo_search.py index 26e7dd467..e72f4594b 100644 --- a/evaluation/scripts/locomo/locomo_search.py +++ b/evaluation/scripts/locomo/locomo_search.py @@ -1,6 +1,24 @@ +import os +import sys +import uuid + + +sys.path.insert( + 0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))) +) +sys.path.insert( + 0, + os.path.join( + os.path.dirname( + os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + ), + "evaluation", + "scripts", + ), +) + import argparse import json -import os from collections import defaultdict from concurrent.futures import ThreadPoolExecutor, as_completed @@ -11,7 +29,8 @@ from dotenv import load_dotenv from mem0 import MemoryClient from tqdm import tqdm -from utils import filter_memory_data +from utils.client import memobase_client, memos_client +from utils.memos_filters import filter_memory_data from zep_cloud.client import Zep from memos.configs.mem_os import MOSConfig @@ -101,6 +120,15 @@ def get_client(frame: str, user_id: str | None = None, version: str = "default", {speaker_2_memories} """ +TEMPLATE_MEMOBASE = """Memories for user {speaker_1_user_id}: + + {speaker_1_memories} + + Memories for user {speaker_2_user_id}: + + {speaker_2_memories} +""" + def mem0_search(client, query, speaker_a_user_id, speaker_b_user_id, top_k=20): start = time() @@ -191,6 +219,38 @@ def memos_search(client, query, conv_id, speaker_a, speaker_b, reversed_client=N return context, duration_ms +def memos_api_search( + client, query, conv_id, speaker_a, speaker_b, top_k, version, reversed_client=None +): + start = time() + speaker_a_user_id = conv_id + "_speaker_a" + search_a_results = client.search( + query=query, user_id=f"{speaker_a_user_id.replace('_', '')}{version}", top_k=top_k + ) + speaker_a_context = "" + for item in search_a_results: + speaker_a_context += f"{item}\n" + + speaker_b_user_id = conv_id + "_speaker_b" + search_b_results = reversed_client.search( + query=query, user_id=f"{speaker_b_user_id.replace('_', '')}{version}", top_k=top_k + ) + speaker_b_context = "" + for item in search_b_results: + speaker_b_context += f"{item}\n" + + context = TEMPLATE_MEMOS.format( + speaker_1=speaker_a, + speaker_1_memories=speaker_a_context, + speaker_2=speaker_b, + speaker_2_memories=speaker_b_context, + ) + + print(query, context) + duration_ms = (time() - start) * 1000 + return context, duration_ms + + def mem0_graph_search(client, query, speaker_a_user_id, speaker_b_user_id, top_k=20): start = time() search_speaker_a_results = client.search( @@ -297,7 +357,58 @@ def zep_search(client, query, group_id, top_k=20): return context, duration_ms -def search_query(client, query, metadata, frame, reversed_client=None, top_k=20): +def memobase_search( + client, query, speaker_a, speaker_b, speaker_a_user_id, speaker_b_user_id, top_k=20 +): + start = time() + speaker_a_memories = memobase_search_memory( + client, speaker_a_user_id, query, max_memory_context_size=top_k * 100 + ) + speaker_b_memories = memobase_search_memory( + client, speaker_b_user_id, query, max_memory_context_size=top_k * 100 + ) + context = TEMPLATE_MEMOBASE.format( + speaker_1_user_id=speaker_a, + speaker_1_memories=speaker_a_memories, + indent=4, + speaker_2_user_id=speaker_b, + speaker_2_memories=speaker_b_memories, + ) + print(query, context) + duration_ms = (time() - start) * 1000 + return (context, duration_ms) + + +def string_to_uuid(s: str, salt="memobase_client") -> str: + return str(uuid.uuid5(uuid.NAMESPACE_DNS, s + salt)) + + +def memobase_search_memory( + client, user_id, query, max_memory_context_size, max_retries=3, retry_delay=1 +): + retries = 0 + real_uid = string_to_uuid(user_id) + u = client.get_user(real_uid, no_get=True) + + while retries < max_retries: + try: + memories = u.context( + max_token_size=max_memory_context_size, + chats=[{"role": "user", "content": query}], + event_similarity_threshold=0.2, + fill_window_with_events=True, + ) + return memories + except Exception as e: + print(f"Error during memory search: {e}") + print("Retrying...") + retries += 1 + if retries >= max_retries: + raise e + time.sleep(retry_delay) + + +def search_query(client, query, metadata, frame, version, reversed_client=None, top_k=20): conv_id = metadata.get("conv_id") speaker_a = metadata.get("speaker_a") speaker_b = metadata.get("speaker_b") @@ -316,7 +427,15 @@ def search_query(client, query, metadata, frame, reversed_client=None, top_k=20) ) elif frame == "memos": context, duration_ms = memos_search( - client, query, conv_id, speaker_a, speaker_b, reversed_client + client, query, conv_id, speaker_a, speaker_b, version, reversed_client + ) + elif frame == "memos-api": + context, duration_ms = memos_api_search( + client, query, conv_id, speaker_a, speaker_b, top_k, version, reversed_client + ) + elif frame == "memobase": + context, duration_ms = memobase_search( + client, query, speaker_a, speaker_b, speaker_a_user_id, speaker_b_user_id, top_k ) return context, duration_ms @@ -364,6 +483,15 @@ def process_user(group_idx, locomo_df, frame, version, top_k=20, num_workers=1): speaker_b_user_id = conv_id + "_speaker_b" client = get_client(frame, speaker_a_user_id, version, top_k=top_k) reversed_client = get_client(frame, speaker_b_user_id, version, top_k=top_k) + elif frame == "memos-api": + speaker_a_user_id = conv_id + "_speaker_a" + speaker_b_user_id = conv_id + "_speaker_b" + client = memos_client(mode="api") + reversed_client = memos_client(mode="api") + client.user_register(user_id=f"{speaker_a_user_id.replace('_', '')}{version}") + reversed_client.user_register(user_id=f"{speaker_b_user_id.replace('_', '')}{version}") + elif frame == "memobase": + client = memobase_client() else: client = get_client(frame, conv_id, version) @@ -372,7 +500,7 @@ def process_qa(qa): if qa.get("category") == 5: return None context, duration_ms = search_query( - client, query, metadata, frame, reversed_client=reversed_client, top_k=top_k + client, query, metadata, frame, version, reversed_client=reversed_client, top_k=top_k ) if not context: @@ -439,8 +567,7 @@ def main(frame, version="default", num_workers=1, top_k=20): parser.add_argument( "--lib", type=str, - choices=["zep", "memos", "mem0", "mem0_graph", "langmem"], - help="Specify the memory framework (zep or memos or mem0 or mem0_graph)", + choices=["zep", "memos", "mem0", "mem0_graph", "memos-api", "memobase"], ) parser.add_argument( "--version", diff --git a/evaluation/scripts/longmemeval/lme_eval.py b/evaluation/scripts/longmemeval/lme_eval.py index 2d54a5acf..384f595be 100644 --- a/evaluation/scripts/longmemeval/lme_eval.py +++ b/evaluation/scripts/longmemeval/lme_eval.py @@ -346,7 +346,7 @@ async def main(frame, version, nlp_options, num_runs=3, num_workers=5): parser.add_argument( "--lib", type=str, - choices=["mem0-local", "mem0-api"], + choices=["mem0-local", "mem0-api", "memos-local", "zep", "memos-api", "zep", "memobase"], ) parser.add_argument( "--version", type=str, default="v1", help="Version of the evaluation framework." diff --git a/evaluation/scripts/longmemeval/lme_ingestion.py b/evaluation/scripts/longmemeval/lme_ingestion.py index aef8e076d..f2df0bd30 100644 --- a/evaluation/scripts/longmemeval/lme_ingestion.py +++ b/evaluation/scripts/longmemeval/lme_ingestion.py @@ -10,7 +10,8 @@ import pandas as pd from tqdm import tqdm -from utils.client import mem0_client, memos_client, zep_client +from utils.client import mem0_client, memobase_client, memos_client, zep_client +from utils.memobase_utils import memobase_add_memory, string_to_uuid from zep_cloud.types import Message @@ -19,7 +20,7 @@ def ingest_session(session, date, user_id, session_id, frame, client): if frame == "zep": for idx, msg in enumerate(session): print( - f"\033[90m[{frame}]\033[0m 💬 Session \033[1;94m{session_id}\033[0m: [\033[93m{idx + 1}/{len(session)}\033[0m] Ingesting message: \033[1m{msg['role']}\033[0m - \033[96m{msg['content'][:50]}...\033[0m at \033[92m{date.isoformat()}\033[0m" + f"\033[90m[{frame}]\033[0m 📝 User \033[1;94m{user_id}\033[0m 💬 Session \033[1;94m{session_id}\033[0m: [\033[93m{idx + 1}/{len(session)}\033[0m] Ingesting message: \033[1m{msg['role']}\033[0m - \033[96m{msg['content'][:50]}...\033[0m at \033[92m{date.isoformat()}\033[0m" ) client.memory.add( session_id=session_id, @@ -53,32 +54,49 @@ def ingest_session(session, date, user_id, session_id, frame, client): print( f"\033[90m[{frame}]\033[0m ✅ Session \033[1;94m{session_id}\033[0m: Ingested \033[93m{len(messages)}\033[0m messages at \033[92m{date.isoformat()}\033[0m" ) - elif frame == "memos-local": + elif frame == "memobase": for idx, msg in enumerate(session): messages.append( { "role": msg["role"], "content": msg["content"][:8000], - "chat_time": date.isoformat(), + "created_at": date.isoformat(), } ) print( - f"\033[90m[{frame}]\033[0m 📝 Session \033[1;94m{session_id}\033[0m: [\033[93m{idx + 1}/{len(session)}\033[0m] Reading message: \033[1m{msg['role']}\033[0m - \033[96m{msg['content'][:50]}...\033[0m at \033[92m{date.isoformat()}\033[0m" + f"\033[90m[{frame}]\033[0m 📝 User \033[1;94m{user_id}\033[0m 💬 Session \033[1;94m{session_id}\033[0m: [\033[93m{idx + 1}/{len(session)}\033[0m] Ingesting message: \033[1m{msg['role']}\033[0m - \033[96m{msg['content'][:50]}...\033[0m at \033[92m{date.isoformat()}\033[0m" + ) + + real_uid = string_to_uuid(user_id) + user = client.get_user(real_uid) + memobase_add_memory(user, messages) + user.flash(sync=True) + print( + f"\033[90m[{frame}]\033[0m ✅ Session \033[1;94m{session_id}\033[0m: Ingested \033[93m{len(messages)}\033[0m messages at \033[92m{date.isoformat()}\033[0m" + ) + elif frame == "memos-local" or frame == "memos-api": + for _idx, msg in enumerate(session): + messages.append( + { + "role": msg["role"], + "content": msg["content"][:8000], + "chat_time": date.isoformat(), + } ) client.add(messages=messages, user_id=user_id) print( f"\033[90m[{frame}]\033[0m ✅ Session \033[1;94m{session_id}\033[0m: Ingested \033[93m{len(messages)}\033[0m messages at \033[92m{date.isoformat()}\033[0m" ) + client.mem_reorganizer_wait() -def ingest_conv(lme_df, version, conv_idx, frame, num_workers=2): +def ingest_conv(lme_df, version, conv_idx, frame): conversation = lme_df.iloc[conv_idx] sessions = conversation["haystack_sessions"] dates = conversation["haystack_dates"] user_id = "lme_exper_user_" + str(conv_idx) - session_id = "lme_exper_session_" + str(conv_idx) print("\n" + "=" * 80) print(f"🔄 \033[1;36mINGESTING CONVERSATION {conv_idx}\033[0m".center(80)) @@ -89,19 +107,10 @@ def ingest_conv(lme_df, version, conv_idx, frame, num_workers=2): print("🔌 \033[1mUsing \033[94mZep client\033[0m \033[1mfor ingestion...\033[0m") # Delete existing user and session if they exist client.user.delete(user_id) - client.memory.delete(session_id) - print( - f"🗑️ Deleted existing user \033[93m{user_id}\033[0m and session \033[93m{session_id}\033[0m from Zep memory..." - ) - # Add user and session to Zep memory + print(f"🗑️ Deleted existing user \033[93m{user_id}\033[0m from Zep memory...") + # Add user to Zep memory client.user.add(user_id=user_id) - client.memory.add_session( - user_id=user_id, - session_id=session_id, - ) - print( - f"➕ Added user \033[93m{user_id}\033[0m and session \033[93m{session_id}\033[0m to Zep memory..." - ) + print(f"➕ Added user \033[93m{user_id}\033[0m to Zep memory...") elif frame == "mem0-local": client = mem0_client(mode="local") print("🔌 \033[1mUsing \033[94mMem0 Local client\033[0m \033[1mfor ingestion...\033[0m") @@ -117,41 +126,49 @@ def ingest_conv(lme_df, version, conv_idx, frame, num_workers=2): elif frame == "memos-local": client = memos_client( mode="local", - db_name=f"lme_{frame}-{version}-{user_id.replace('_', '')}", + db_name=f"lme_{frame}-{version}", user_id=user_id, top_k=20, mem_cube_path=f"results/lme/{frame}-{version}/storages/{user_id}", - mem_cube_config_path="configs/mem_cube_config.json", + mem_cube_config_path="configs/mu_mem_cube_config.json", mem_os_config_path="configs/mos_memos_config.json", addorsearch="add", ) print("🔌 \033[1mUsing \033[94mMemos Local client\033[0m \033[1mfor ingestion...\033[0m") - - with ThreadPoolExecutor(max_workers=num_workers) as executor: - futures = [] - - for idx, session in enumerate(sessions): - date = dates[idx] + " UTC" - date_format = "%Y/%m/%d (%a) %H:%M UTC" - date_string = datetime.strptime(date, date_format).replace(tzinfo=timezone.utc) - - future = executor.submit( - ingest_session, session, date_string, user_id, session_id, frame, client + elif frame == "memos-api": + client = memos_client(mode="api") + elif frame == "memobase": + client = memobase_client() + print("🔌 \033[1mUsing \033[94mMemobase client\033[0m \033[1mfor ingestion...\033[0m") + client.delete_user(string_to_uuid(user_id)) + print(f"🗑️ Deleted existing user \033[93m{user_id}\033[0m from Memobase memory...") + + for idx, session in enumerate(sessions): + session_id = user_id + "_lme_exper_session_" + str(idx) + if frame == "zep": + client.memory.add_session( + user_id=user_id, + session_id=session_id, + ) + print( + f"➕ Added session \033[93m{session_id}\033[0m for user \033[93m{user_id}\033[0m to Zep memory..." ) - futures.append(future) - if len(session) == 0: - print(f"\033[93m⚠️ Skipping empty session {idx} in conversation {conv_idx}\033[0m") - continue + if len(session) == 0: + print(f"\033[93m⚠️ Skipping empty session {idx} in conversation {conv_idx}\033[0m") + continue - for future in tqdm( - as_completed(futures), total=len(futures), desc=f"📊 Ingesting user {conv_idx}" - ): - try: - future.result() - except Exception as e: - print(f"\033[91m❌ Error ingesting session: {e}\033[0m") + date = dates[idx] + " UTC" + date_format = "%Y/%m/%d (%a) %H:%M UTC" + date_string = datetime.strptime(date, date_format).replace(tzinfo=timezone.utc) + + try: + ingest_session(session, date_string, user_id, session_id, frame, client) + except Exception as e: + print(f"\033[91m❌ Error ingesting session: {e}\033[0m") + if frame == "memos-local": + client.mem_reorganizer_off() print("=" * 80) @@ -170,8 +187,21 @@ def main(frame, version, num_workers=2): print("-" * 80) start_time = datetime.now() - for session_idx in range(num_multi_sessions): - ingest_conv(lme_df, version, session_idx, frame, num_workers=num_workers) + + with ThreadPoolExecutor(max_workers=num_workers) as executor: + futures = [] + for session_idx in range(num_multi_sessions): + future = executor.submit(ingest_conv, lme_df, version, session_idx, frame) + futures.append(future) + + for future in tqdm( + as_completed(futures), total=len(futures), desc="📊 Processing conversations" + ): + try: + future.result() + except Exception as e: + print(f"\033[91m❌ Error processing conversation: {e}\033[0m") + end_time = datetime.now() elapsed_time = end_time - start_time elapsed_time_str = str(elapsed_time).split(".")[0] @@ -193,7 +223,7 @@ def main(frame, version, num_workers=2): parser.add_argument( "--lib", type=str, - choices=["mem0-local", "mem0-api", "memos-local"], + choices=["mem0-local", "mem0-api", "memos-local", "memos-api", "zep", "memobase"], ) parser.add_argument( "--version", type=str, default="v1", help="Version of the evaluation framework." diff --git a/evaluation/scripts/longmemeval/lme_metric.py b/evaluation/scripts/longmemeval/lme_metric.py index be285123f..69f7748e0 100644 --- a/evaluation/scripts/longmemeval/lme_metric.py +++ b/evaluation/scripts/longmemeval/lme_metric.py @@ -258,7 +258,7 @@ def calculate_scores(data, grade_path, output_path): parser.add_argument( "--lib", type=str, - choices=["mem0-local", "mem0-api"], + choices=["mem0-local", "mem0-api", "memos-local", "memos-api", "zep", "memobase"], ) parser.add_argument( "--version", type=str, default="v1", help="Version of the evaluation framework." diff --git a/evaluation/scripts/longmemeval/lme_responses.py b/evaluation/scripts/longmemeval/lme_responses.py index 9d5f8c1ab..e1e341826 100644 --- a/evaluation/scripts/longmemeval/lme_responses.py +++ b/evaluation/scripts/longmemeval/lme_responses.py @@ -145,7 +145,7 @@ def main(frame, version, num_workers=4): parser.add_argument( "--lib", type=str, - choices=["mem0-local", "mem0-api"], + choices=["mem0-local", "mem0-api", "memos-local", "memos-api", "zep", "memobase"], ) parser.add_argument( "--version", type=str, default="v1", help="Version of the evaluation framework." diff --git a/evaluation/scripts/longmemeval/lme_search.py b/evaluation/scripts/longmemeval/lme_search.py index 0643c07ff..898ab7e27 100644 --- a/evaluation/scripts/longmemeval/lme_search.py +++ b/evaluation/scripts/longmemeval/lme_search.py @@ -13,11 +13,13 @@ import pandas as pd from tqdm import tqdm -from utils.client import mem0_client, memos_client, zep_client +from utils.client import mem0_client, memobase_client, memos_client, zep_client +from utils.memobase_utils import memobase_search_memory from utils.memos_filters import filter_memory_data from utils.prompts import ( MEM0_CONTEXT_TEMPLATE, MEM0_GRAPH_CONTEXT_TEMPLATE, + MEMOBASE_CONTEXT_TEMPLATE, MEMOS_CONTEXT_TEMPLATE, ZEP_CONTEXT_TEMPLATE, ) @@ -111,21 +113,37 @@ def mem0_search(client, user_id, query, top_k=20, enable_graph=False, frame="mem return context, duration_ms -def memos_search(client, user_id, query, frame="memos-local"): +def memos_search(client, user_id, query, top_k, frame="memos-local"): start = time() + if frame == "memos-local": + results = client.search( + query=query, + user_id=user_id, + ) - results = client.search( - query=query, - user_id=user_id, - ) + results = filter_memory_data(results)["text_mem"][0]["memories"] + search_memories = "\n".join([f" - {item['memory']}" for item in results]) - search_memories = filter_memory_data(results)["text_mem"][0]["memories"] + elif frame == "memos-api": + results = client.search(query=query, user_id=user_id, top_k=top_k) + search_memories = "\n".join([f" - {item}" for item in results]) context = MEMOS_CONTEXT_TEMPLATE.format(user_id=user_id, memories=search_memories) duration_ms = (time() - start) * 1000 return context, duration_ms +def memobase_search(client, user_id, query, top_k=20): + start = time() + memories = memobase_search_memory(client, user_id, query, max_memory_context_size=top_k * 100) + context = MEMOBASE_CONTEXT_TEMPLATE.format( + user_id=user_id, + memories=memories, + ) + duration_ms = (time() - start) * 1000 + return context, duration_ms + + def process_user(lme_df, conv_idx, frame, version, top_k=20): row = lme_df.iloc[conv_idx] question = row["question"] @@ -175,17 +193,27 @@ def process_user(lme_df, conv_idx, frame, version, top_k=20): elif frame == "memos-local": client = memos_client( mode="local", - db_name=f"lme_{frame}-{version}-{user_id.replace('_', '')}", + db_name=f"lme_{frame}-{version}", user_id=user_id, - top_k=20, + top_k=top_k, mem_cube_path=f"results/lme/{frame}-{version}/storages/{user_id}", - mem_cube_config_path="configs/mem_cube_config.json", + mem_cube_config_path="configs/mu_mem_cube_config.json", mem_os_config_path="configs/mos_memos_config.json", addorsearch="search", ) print("🔌 \033[1mUsing \033[94mMemos Local client\033[0m \033[1mfor search...\033[0m") context, duration_ms = memos_search(client, user_id, question, frame=frame) + elif frame == "memobase": + client = memobase_client() + print("🔌 \033[1mUsing \033[94mMemobase client\033[0m \033[1mfor search...\033[0m") + context, duration_ms = memobase_search_memory(client, user_id, question, top_k=top_k) + elif frame == "memos-api": + client = memos_client( + mode="api", + ) + print("🔌 \033[1mUsing \033[94mMemos API client\033[0m \033[1mfor search...\033[0m") + context, duration_ms = memos_search(client, user_id, question, top_k=top_k, frame=frame) search_results[user_id].append( { "question": question, @@ -282,7 +310,11 @@ def main(frame, version, top_k=20, num_workers=2): if __name__ == "__main__": parser = argparse.ArgumentParser(description="LongMemeval Search Script") - parser.add_argument("--lib", type=str, choices=["mem0-local", "mem0-api", "memos-local"]) + parser.add_argument( + "--lib", + type=str, + choices=["mem0-local", "mem0-api", "memos-local", "memos-api", "zep", "memobase"], + ) parser.add_argument( "--version", type=str, default="v1", help="Version of the evaluation framework." ) diff --git a/evaluation/scripts/run_lme_eval.sh b/evaluation/scripts/run_lme_eval.sh index 96e430fd6..fc9031ea0 100755 --- a/evaluation/scripts/run_lme_eval.sh +++ b/evaluation/scripts/run_lme_eval.sh @@ -2,8 +2,8 @@ # Common parameters for all scripts LIB="memos-local" -VERSION="071503" -WORKERS=10 +VERSION="072202" +WORKERS=50 TOPK=20 echo "Running lme_ingestion.py..." @@ -20,4 +20,25 @@ if [ $? -ne 0 ]; then exit 1 fi +echo "Running lme_responses.py..." +CUDA_VISIBLE_DEVICES=0 python scripts/longmemeval/lme_responses.py --lib $LIB --version $VERSION --workers $WORKERS +if [ $? -ne 0 ]; then + echo "Error running lme_responses.py" + exit 1 +fi + +echo "Running lme_eval.py..." +CUDA_VISIBLE_DEVICES=0 python scripts/longmemeval/lme_eval.py --lib $LIB --version $VERSION --workers $WORKERS +if [ $? -ne 0 ]; then + echo "Error running lme_eval.py" + exit 1 +fi + +echo "Running lme_metric.py..." +CUDA_VISIBLE_DEVICES=0 python scripts/longmemeval/lme_metric.py --lib $LIB --version $VERSION +if [ $? -ne 0 ]; then + echo "Error running lme_metric.py" + exit 1 +fi + echo "All scripts completed successfully!" diff --git a/evaluation/scripts/run_locomo_eval.sh b/evaluation/scripts/run_locomo_eval.sh index df1a865f2..89a09729c 100755 --- a/evaluation/scripts/run_locomo_eval.sh +++ b/evaluation/scripts/run_locomo_eval.sh @@ -1,17 +1,17 @@ #!/bin/bash # Common parameters for all scripts -LIB="memos" -VERSION="063001" +LIB="memos-api" +VERSION="072001" WORKERS=10 TOPK=20 -echo "Running locomo_ingestion.py..." -CUDA_VISIBLE_DEVICES=0 python scripts/locomo/locomo_ingestion.py --lib $LIB --version $VERSION --workers $WORKERS -if [ $? -ne 0 ]; then - echo "Error running locomo_ingestion.py" - exit 1 -fi +# echo "Running locomo_ingestion.py..." +# CUDA_VISIBLE_DEVICES=0 python scripts/locomo/locomo_ingestion.py --lib $LIB --version $VERSION --workers $WORKERS +# if [ $? -ne 0 ]; then +# echo "Error running locomo_ingestion.py" +# exit 1 +# fi echo "Running locomo_search.py..." CUDA_VISIBLE_DEVICES=0 python scripts/locomo/locomo_search.py --lib $LIB --version $VERSION --top_k $TOPK --workers $WORKERS diff --git a/evaluation/scripts/run_pm_eval.sh b/evaluation/scripts/run_pm_eval.sh new file mode 100755 index 000000000..a3321cdf9 --- /dev/null +++ b/evaluation/scripts/run_pm_eval.sh @@ -0,0 +1,37 @@ +#!/bin/bash + +# Common parameters for all scripts +LIB="memos-local" +VERSION="072201" +WORKERS=10 +TOPK=20 + +echo "Running pm_ingestion.py..." +CUDA_VISIBLE_DEVICES=0 python scripts/personamem/pm_ingestion.py --lib $LIB --version $VERSION --workers $WORKERS +if [ $? -ne 0 ]; then + echo "Error running pm_ingestion.py" + exit 1 +fi + +echo "Running pm_search.py..." +CUDA_VISIBLE_DEVICES=0 python scripts/personamem/pm_search.py --lib $LIB --version $VERSION --top_k $TOPK --workers $WORKERS +if [ $? -ne 0 ]; then + echo "Error running pm_search.py" + exit 1 +fi + +echo "Running pm_responses.py..." +CUDA_VISIBLE_DEVICES=0 python scripts/personamem/pm_responses.py --lib $LIB --version $VERSION --workers $WORKERS +if [ $? -ne 0 ]; then + echo "Error running pm_responses.py" + exit 1 +fi + +echo "Running pm_metric.py..." +CUDA_VISIBLE_DEVICES=0 python scripts/personamem/pm_metric.py --lib $LIB --version $VERSION +if [ $? -ne 0 ]; then + echo "Error running pm_metric.py" + exit 1 +fi + +echo "All scripts completed successfully!" diff --git a/evaluation/scripts/utils/client.py b/evaluation/scripts/utils/client.py index ddb144f6f..33aea7497 100644 --- a/evaluation/scripts/utils/client.py +++ b/evaluation/scripts/utils/client.py @@ -4,16 +4,20 @@ from dotenv import load_dotenv from mem0 import MemoryClient +from memobase import MemoBaseClient from zep_cloud.client import Zep from zep_cloud.types import Message sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from memobase import ChatBlob + 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 +from memos.mem_os.product import MOSProduct from utils.mem0_local import Mem0Client +from utils.memos_api import MemOSAPI from utils.memos_filters import filter_memory_data @@ -57,7 +61,7 @@ def memos_client( mos_config_data = json.load(f) mos_config_data["top_k"] = top_k mos_config = MOSConfig(**mos_config_data) - memos = MOS(mos_config) + memos = MOSProduct(mos_config) memos.create_user(user_id=user_id) if addorsearch == "add": @@ -68,24 +72,35 @@ def memos_client( mem_cube_config_data["text_mem"]["config"]["graph_db"]["config"]["db_name"] = ( f"{db_name.replace('_', '')}" ) + mem_cube_config_data["text_mem"]["config"]["graph_db"]["config"]["user_name"] = user_id + mem_cube_config_data["text_mem"]["config"]["reorganize"] = True mem_cube_config = GeneralMemCubeConfig.model_validate(mem_cube_config_data) mem_cube = GeneralMemCube(mem_cube_config) if not os.path.exists(mem_cube_path): mem_cube.dump(mem_cube_path) - memos.register_mem_cube( - mem_cube_name_or_path=mem_cube_path, - mem_cube_id=user_id, + memos.user_register( user_id=user_id, + user_name=user_id, + interests=f"I'm {user_id}", + default_mem_cube=mem_cube, ) elif mode == "api": - pass + memos = MemOSAPI(base_url=os.getenv("MEMOS_BASE_URL")) return memos +def memobase_client(): + client = MemoBaseClient( + project_url=os.getenv("MEMOBASE_PROJECT_URL"), + api_key=os.getenv("MEMOBASE_API_KEY"), + ) + return client + + if __name__ == "__main__": # Example usage of the Zep client zep = zep_client() @@ -196,3 +211,25 @@ def memos_client( search_result_b = memos_b.search(query="football", user_id="alice") filtered_search_result_b = filter_memory_data(search_result_b)["text_mem"][0]["memories"] print("Search results in Memos B:", filtered_search_result_b) + + # Example usage of MemoBase client + client = memobase_client() + print("MemoBase client initialized successfully.") + + # Example of adding a user and retrieving user information + user_id = client.add_user() + user = client.get_user(user_id) + + # Example of adding a chat blob to the user + print(f"Adding chat blob for user {user_id}...") + b = ChatBlob( + messages=[ + {"role": "user", "content": "Hi, I'm here again"}, + {"role": "assistant", "content": "Hi, Gus! How can I help you?"}, + ] + ) + bid = user.insert(b) + + # Example of retrieving the context of the user + context = user.context() + print(context) diff --git a/evaluation/scripts/utils/memobase_utils.py b/evaluation/scripts/utils/memobase_utils.py new file mode 100644 index 000000000..dcf06ea31 --- /dev/null +++ b/evaluation/scripts/utils/memobase_utils.py @@ -0,0 +1,46 @@ +import time +import uuid + +from memobase import ChatBlob + + +def string_to_uuid(s: str, salt="memobase_client") -> str: + return str(uuid.uuid5(uuid.NAMESPACE_DNS, s + salt)) + + +def memobase_add_memory(user, message, retries=3): + for attempt in range(retries): + try: + _ = user.insert(ChatBlob(messages=message), sync=True) + return + except Exception as e: + if attempt < retries - 1: + time.sleep(1) + continue + else: + raise e + + +def memobase_search_memory( + client, user_id, query, max_memory_context_size, max_retries=3, retry_delay=1 +): + retries = 0 + real_uid = string_to_uuid(user_id) + u = client.get_user(real_uid, no_get=True) + + while retries < max_retries: + try: + memories = u.context( + max_token_size=max_memory_context_size, + chats=[{"role": "user", "content": query}], + event_similarity_threshold=0.2, + fill_window_with_events=True, + ) + return memories + except Exception as e: + print(f"Error during memory search: {e}") + print("Retrying...") + retries += 1 + if retries >= max_retries: + raise e + time.sleep(retry_delay) diff --git a/evaluation/scripts/utils/memos_api.py b/evaluation/scripts/utils/memos_api.py new file mode 100644 index 000000000..7b7f2a061 --- /dev/null +++ b/evaluation/scripts/utils/memos_api.py @@ -0,0 +1,63 @@ +import json + +import requests + + +class MemOSAPI: + def __init__(self, base_url: str = "http://localhost:8000"): + self.base_url = base_url + self.headers = {"Content-Type": "application/json"} + + def user_register(self, user_id: str): + """Register a user.""" + url = f"{self.base_url}/users/register" + payload = json.dumps({"user_id": user_id}) + response = requests.request("POST", url, data=payload, headers=self.headers) + return response.text + + def add(self, messages: list[dict], user_id: str | None = None): + """Create memories.""" + register_res = json.loads(self.user_register(user_id)) + cube_id = register_res["data"]["mem_cube_id"] + url = f"{self.base_url}/add" + payload = json.dumps({"messages": messages, "user_id": user_id, "mem_cube_id": cube_id}) + + response = requests.request("POST", url, data=payload, headers=self.headers) + return response.text + + def search(self, query: str, user_id: str | None = None, top_k: int = 10): + """Search memories.""" + url = f"{self.base_url}/search" + payload = json.dumps( + { + "query": query, + "user_id": user_id, + } + ) + + response = requests.request("POST", url, data=payload, headers=self.headers) + if response.status_code != 200: + response.raise_for_status() + else: + result = json.loads(response.text)["data"]["text_mem"][0]["memories"] + text_memories = [item["memory"] for item in result][:top_k] + return text_memories + + +if __name__ == "__main__": + client = MemOSAPI(base_url="http://localhost:8000") + # Example usage + try: + messages = [ + { + "role": "user", + "content": "I went to the store and bought a red apple.", + "chat_time": "2023-10-01T12:00:00Z", + } + ] + add_response = client.add(messages, user_id="user789") + print("Add memory response:", add_response) + search_response = client.search("red apple", user_id="user789", top_k=1) + print("Search memory response:", search_response) + except requests.RequestException as e: + print("An error occurred:", e) diff --git a/evaluation/scripts/utils/prompts.py b/evaluation/scripts/utils/prompts.py index 6515619ec..dd83acdc2 100644 --- a/evaluation/scripts/utils/prompts.py +++ b/evaluation/scripts/utils/prompts.py @@ -28,6 +28,37 @@ Answer: """ +PM_ANSWER_PROMPT = """ + You are a helpful assistant tasked with selecting the best answer to a user question, based solely on summarized conversation memories. + + # CONTEXT: + The following are summarized facts and preferences extracted from prior user conversations. Use only these memories to answer the question. + + {context} + + # INSTRUCTIONS: + 1. Carefully read and reason over the memory summary. + 2. Evaluate each of the four answer choices (a) through (d). + 3. Choose the single best-supported answer based on the information in memory. + 4. Output ONLY the final choice in the format (a), (b), (c), or (d), placed directly after the token . + + # IMPORTANT RULES: + - Your final answer **must appear after** the token . + - Your final answer **must use parentheses**, like (a) or (b). + - Do NOT list multiple choices. Choose only one. + - Do NOT include extra text after . Just output the answer. + + # QUESTION: + {question} + + # OPTIONS: + {options} + + Final Answer: + +""" + + ZEP_CONTEXT_TEMPLATE = """ FACTS and ENTITIES represent relevant context to the current conversation. @@ -53,6 +84,12 @@ {memories} """ +MEMOBASE_CONTEXT_TEMPLATE = """ + Memories for user {user_id}: + + {memories} +""" + MEM0_GRAPH_CONTEXT_TEMPLATE = """ Memories for user {user_id}: diff --git a/examples/basic_modules/nebular_example.py b/examples/basic_modules/nebular_example.py index 0cdddea6e..8abbbcbde 100644 --- a/examples/basic_modules/nebular_example.py +++ b/examples/basic_modules/nebular_example.py @@ -16,6 +16,20 @@ load_dotenv() + +def show(nebular_data): + from memos.configs.graph_db import Neo4jGraphDBConfig + from memos.graph_dbs.neo4j import Neo4jGraphDB + + tree_config = Neo4jGraphDBConfig.from_json_file("../../examples/data/config/neo4j_config.json") + tree_config.use_multi_db = True + tree_config.db_name = "nebular-show" + + neo4j_db = Neo4jGraphDB(tree_config) + neo4j_db.clear() + neo4j_db.import_graph(nebular_data) + + embedder_config = EmbedderConfigFactory.model_validate( { "backend": "universal_api", @@ -42,13 +56,13 @@ def example_multi_db(db_name: str = "paper"): config = GraphDBConfigFactory( backend="nebular", config={ - "hosts": json.loads(os.getenv("NEBULAR_HOSTS", "localhost")), - "user_name": os.getenv("NEBULAR_USER", "root"), + "uri": json.loads(os.getenv("NEBULAR_HOSTS", "localhost")), + "user": os.getenv("NEBULAR_USER", "root"), "password": os.getenv("NEBULAR_PASSWORD", "xxxxxx"), "space": db_name, + "use_multi_db": True, "auto_create": True, "embedding_dimension": 3072, - "use_multi_db": True, }, ) @@ -93,20 +107,21 @@ def example_shared_db(db_name: str = "shared-traval-group"): Multiple users' data in the same Neo4j DB with user_name as a tag. """ # users - user_list = ["root"] + user_list = ["travel_member_alice", "travel_member_bob"] for user_name in user_list: # Step 1: Build factory config config = GraphDBConfigFactory( backend="nebular", config={ - "hosts": json.loads(os.getenv("NEBULAR_HOSTS", "localhost")), - "user_name": os.getenv("NEBULAR_USER", "root"), + "uri": json.loads(os.getenv("NEBULAR_HOSTS", "localhost")), + "user": os.getenv("NEBULAR_USER", "root"), "password": os.getenv("NEBULAR_PASSWORD", "xxxxxx"), "space": db_name, + "user_name": user_name, + "use_multi_db": False, "auto_create": True, "embedding_dimension": 3072, - "use_multi_db": False, }, ) @@ -187,10 +202,11 @@ def example_shared_db(db_name: str = "shared-traval-group"): config_alice = GraphDBConfigFactory( backend="nebular", config={ - "hosts": json.loads(os.getenv("NEBULAR_HOSTS", "localhost")), - "user_name": os.getenv("NEBULAR_USER", "root"), + "uri": json.loads(os.getenv("NEBULAR_HOSTS", "localhost")), + "user": os.getenv("NEBULAR_USER", "root"), "password": os.getenv("NEBULAR_PASSWORD", "xxxxxx"), "space": db_name, + "user_name": user_list[0], "auto_create": True, "embedding_dimension": 3072, "use_multi_db": False, @@ -215,13 +231,14 @@ def run_user_session( config = GraphDBConfigFactory( backend="nebular", config={ - "hosts": json.loads(os.getenv("NEBULAR_HOSTS", "localhost")), - "user_name": os.getenv("NEBULAR_USER", "root"), + "uri": json.loads(os.getenv("NEBULAR_HOSTS", "localhost")), + "user": os.getenv("NEBULAR_USER", "root"), "password": os.getenv("NEBULAR_PASSWORD", "xxxxxx"), "space": db_name, + "user_name": user_name, + "use_multi_db": False, "auto_create": True, "embedding_dimension": 3072, - "use_multi_db": False, }, ) graph = GraphStoreFactory.from_config(config) @@ -242,6 +259,7 @@ def run_user_session( memory_time="2024-01-01", status="activated", visibility="public", + tags=["research", "rl"], updated_at=now, embedding=embed_memory_item(topic_text), ), @@ -300,7 +318,6 @@ def run_user_session( print("🔍 Search result:", node["memory"]) # === Step 5: Tag-based neighborhood discovery === - # TODO neighbors = graph.get_neighbors_by_tag(["concept"], exclude_ids=[], top_k=2) print("📎 Tag-related nodes:", [neighbor["memory"] for neighbor in neighbors]) @@ -319,7 +336,7 @@ def run_user_session( graph.update_node( concept_items[0].id, {"confidence": 99.0, "created_at": "2025-07-24T20:11:56.375687"} ) - graph.remove_oldest_memory("LongTermMemory", keep_latest=3) + 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) @@ -328,6 +345,30 @@ def run_user_session( graph.import_graph(exported) print("📦 Graph exported and re-imported, total nodes:", len(exported["nodes"])) + # ==================================== + # 🔍 Step 10: extra function + # ==================================== + print(f"\n=== 🔍 Extra Tests for user: {user_name} ===") + + print(" - Memory count:", graph.get_memory_count("LongTermMemory")) + print(" - Node count:", graph.count_nodes("LongTermMemory")) + print(" - All LongTermMemory items:", graph.get_all_memory_items("LongTermMemory")) + + if len(exported["edges"]) > 0: + n1, n2 = exported["edges"][0]["source"], exported["edges"][0]["target"] + print(" - Edge exists?", graph.edge_exists(n1, n2, exported["edges"][0]["type"])) + print(" - Edges for node:", graph.get_edges(n1)) + + filters = [{"field": "memory_type", "op": "=", "value": "LongTermMemory"}] + print(" - Metadata query result:", graph.get_by_metadata(filters)) + print( + " - Optimization candidates:", graph.get_structure_optimization_candidates("LongTermMemory") + ) + try: + graph.drop_database() + except ValueError as e: + print(" - drop_database raised ValueError as expected:", e) + def example_complex_shared_db(db_name: str = "shared-traval-group-complex"): # User 1: Alice explores structured memory for LLMs @@ -369,4 +410,4 @@ def example_complex_shared_db(db_name: str = "shared-traval-group-complex"): 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") + example_complex_shared_db(db_name="shared-traval-group-complex-new11") diff --git a/examples/basic_modules/neo4j_example.py b/examples/basic_modules/neo4j_example.py index 082ad8c3e..ea68975cc 100644 --- a/examples/basic_modules/neo4j_example.py +++ b/examples/basic_modules/neo4j_example.py @@ -1,3 +1,5 @@ +import os + from datetime import datetime from memos.configs.embedder import EmbedderConfigFactory @@ -8,7 +10,15 @@ embedder_config = EmbedderConfigFactory.model_validate( - {"backend": "ollama", "config": {"model_name_or_path": "nomic-embed-text:latest"}} + { + "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"), + }, + } ) embedder = EmbedderFactory.from_config(embedder_config) @@ -27,7 +37,7 @@ def example_multi_db(db_name: str = "paper"): "password": "12345678", "db_name": db_name, "auto_create": True, - "embedding_dimension": 768, + "embedding_dimension": 3072, "use_multi_db": True, }, ) @@ -268,7 +278,7 @@ def example_shared_db(db_name: str = "shared-traval-group"): "user_name": user_name, "use_multi_db": False, "auto_create": True, - "embedding_dimension": 768, + "embedding_dimension": 3072, }, ) # Step 2: Instantiate graph store @@ -331,7 +341,7 @@ def example_shared_db(db_name: str = "shared-traval-group"): "password": "12345678", "db_name": db_name, "user_name": user_list[0], - "embedding_dimension": 768, + "embedding_dimension": 3072, }, ) graph_alice = GraphStoreFactory.from_config(config_alice) @@ -362,14 +372,14 @@ def run_user_session( "user_name": user_name, "use_multi_db": False, "auto_create": False, # Neo4j Community does not allow auto DB creation - "embedding_dimension": 768, + "embedding_dimension": 3072, "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, + "vector_dimension": 3072, "distance_metric": "cosine", "host": "localhost", "port": 6333, @@ -388,7 +398,7 @@ def run_user_session( "user_name": user_name, "use_multi_db": False, "auto_create": True, - "embedding_dimension": 768, + "embedding_dimension": 3072, }, ) graph = GraphStoreFactory.from_config(config) diff --git a/examples/core_memories/general_textual_memory.py b/examples/core_memories/general_textual_memory.py index 2ecbc7826..f71e2ef2e 100644 --- a/examples/core_memories/general_textual_memory.py +++ b/examples/core_memories/general_textual_memory.py @@ -1,6 +1,7 @@ from memos.configs.memory import MemoryConfigFactory from memos.memories.factory import MemoryFactory + config = MemoryConfigFactory( backend="general_text", config={ diff --git a/examples/core_memories/tree_textual_memory.py b/examples/core_memories/tree_textual_memory.py index 17f68832e..47dc51e41 100644 --- a/examples/core_memories/tree_textual_memory.py +++ b/examples/core_memories/tree_textual_memory.py @@ -203,6 +203,18 @@ def embed_memory_item(memory: str) -> list[float]: print(f"{i}'th similar result is: " + str(r["memory"])) print(f"Successfully search {len(results)} memories") +# try this when use 'fine' mode (Note that you should pass the internet Config, refer to examples/core_memories/textual_internet_memoy.py) +results_fine_search = my_tree_textual_memory.search( + "Recent news in NewYork", + top_k=10, + mode="fine", + info={"query": "Recent news in NewYork", "user_id": "111", "session": "2234"}, +) +for i, r in enumerate(results_fine_search): + r = r.to_dict() + print(f"{i}'th similar result is: " + str(r["memory"])) +print(f"Successfully search {len(results_fine_search)} memories") + # find related nodes related_nodes = my_tree_textual_memory.get_relevant_subgraph("Painting") @@ -235,7 +247,6 @@ def embed_memory_item(memory: str) -> list[float]: # close the synchronous thread in memory manager my_tree_textual_memory.memory_manager.close() - # my_tree_textual_memory.dump my_tree_textual_memory.dump("tmp/my_tree_textual_memory") my_tree_textual_memory.drop() diff --git a/examples/data/config/mem_scheduler/memos_config_w_scheduler_and_openai.yaml b/examples/data/config/mem_scheduler/memos_config_w_scheduler_and_openai.yaml index 1da4dad13..b329e0fc4 100644 --- a/examples/data/config/mem_scheduler/memos_config_w_scheduler_and_openai.yaml +++ b/examples/data/config/mem_scheduler/memos_config_w_scheduler_and_openai.yaml @@ -41,6 +41,7 @@ mem_scheduler: thread_pool_max_workers: 10 consume_interval_seconds: 1 enable_parallel_dispatch: true + enable_act_memory_update: false max_turns_window: 20 top_k: 5 enable_textual_memory: true diff --git a/examples/mem_scheduler/memos_w_scheduler.py b/examples/mem_scheduler/memos_w_scheduler.py index c00845d57..d67b6715a 100644 --- a/examples/mem_scheduler/memos_w_scheduler.py +++ b/examples/mem_scheduler/memos_w_scheduler.py @@ -1,25 +1,17 @@ import shutil import sys -from datetime import datetime from pathlib import Path from queue import Queue from typing import TYPE_CHECKING from memos.configs.mem_cube import GeneralMemCubeConfig from memos.configs.mem_os import MOSConfig -from memos.configs.mem_scheduler import AuthConfig, SchedulerConfigFactory +from memos.configs.mem_scheduler import AuthConfig from memos.log import get_logger from memos.mem_cube.general import GeneralMemCube from memos.mem_os.main import MOS from memos.mem_scheduler.general_scheduler import GeneralScheduler -from memos.mem_scheduler.scheduler_factory import SchedulerFactory -from memos.mem_scheduler.schemas.general_schemas import ( - ANSWER_LABEL, - QUERY_LABEL, -) -from memos.mem_scheduler.schemas.message_schemas import ScheduleMessageItem -from memos.mem_scheduler.utils.misc_utils import parse_yaml if TYPE_CHECKING: @@ -78,122 +70,56 @@ def init_task(): return conversations, questions -def run_with_automatic_scheduler_init(): +def run_with_scheduler_init(): print("==== run_with_automatic_scheduler_init ====") conversations, questions = init_task() - config = parse_yaml( - f"{BASE_DIR}/examples/data/config/mem_scheduler/memos_config_w_scheduler.yaml" + # set configs + mos_config = MOSConfig.from_yaml_file( + f"{BASE_DIR}/examples/data/config/mem_scheduler/memos_config_w_scheduler_and_openai.yaml" ) - mos_config = MOSConfig(**config) - mos = MOS(mos_config) - - user_id = "user_1" - mos.create_user(user_id) - - config = GeneralMemCubeConfig.from_yaml_file( + mem_cube_config = GeneralMemCubeConfig.from_yaml_file( f"{BASE_DIR}/examples/data/config/mem_scheduler/mem_cube_config.yaml" ) - mem_cube_id = "mem_cube_5" - mem_cube_name_or_path = f"{BASE_DIR}/outputs/mem_scheduler/{user_id}/{mem_cube_id}" - if Path(mem_cube_name_or_path).exists(): - shutil.rmtree(mem_cube_name_or_path) - print(f"{mem_cube_name_or_path} is not empty, and has been removed.") # default local graphdb uri if AuthConfig.default_config_exists(): auth_config = AuthConfig.from_local_yaml() - config.text_mem.config.graph_db.config.uri = auth_config.graph_db.uri - mem_cube = GeneralMemCube(config) - mem_cube.dump(mem_cube_name_or_path) - mos.register_mem_cube( - mem_cube_name_or_path=mem_cube_name_or_path, mem_cube_id=mem_cube_id, user_id=user_id - ) - mos.add(conversations, user_id=user_id, mem_cube_id=mem_cube_id) + mos_config.mem_reader.config.llm.config.api_key = auth_config.openai.api_key + mos_config.mem_reader.config.llm.config.api_base = auth_config.openai.base_url - for item in questions: - query = item["question"] - response = mos.chat(query, user_id=user_id) - print(f"Query:\n {query}\n\nAnswer:\n {response}") + mem_cube_config.text_mem.config.graph_db.config.uri = auth_config.graph_db.uri - show_web_logs(mem_scheduler=mos.mem_scheduler) - - mos.mem_scheduler.stop() - - -def run_with_manual_scheduler_init(): - print("==== run_with_manual_scheduler_init ====") - conversations, questions = init_task() - - config = parse_yaml( - f"{BASE_DIR}/examples/data/config/mem_scheduler/memos_config_wo_scheduler.yaml" - ) - - mos_config = MOSConfig(**config) + # Initialization mos = MOS(mos_config) user_id = "user_1" mos.create_user(user_id) - config = GeneralMemCubeConfig.from_yaml_file( - f"{BASE_DIR}/examples/data/config/mem_scheduler/mem_cube_config.yaml" - ) mem_cube_id = "mem_cube_5" mem_cube_name_or_path = f"{BASE_DIR}/outputs/mem_scheduler/{user_id}/{mem_cube_id}" + if Path(mem_cube_name_or_path).exists(): shutil.rmtree(mem_cube_name_or_path) print(f"{mem_cube_name_or_path} is not empty, and has been removed.") - # default local graphdb uri - if AuthConfig.default_config_exists(): - auth_config = AuthConfig.from_local_yaml() - config.text_mem.config.graph_db.config.uri = auth_config.graph_db.uri - - mem_cube = GeneralMemCube(config) + mem_cube = GeneralMemCube(mem_cube_config) mem_cube.dump(mem_cube_name_or_path) mos.register_mem_cube( mem_cube_name_or_path=mem_cube_name_or_path, mem_cube_id=mem_cube_id, user_id=user_id ) - example_scheduler_config_path = ( - f"{BASE_DIR}/examples/data/config/mem_scheduler/general_scheduler_config.yaml" - ) - scheduler_config = SchedulerConfigFactory.from_yaml_file( - yaml_path=example_scheduler_config_path - ) - mem_scheduler = SchedulerFactory.from_config(scheduler_config) - mem_scheduler.initialize_modules(chat_llm=mos.chat_llm) - - mos.mem_scheduler = mem_scheduler - - mos.mem_scheduler.start() - mos.add(conversations, user_id=user_id, mem_cube_id=mem_cube_id) for item in questions: + print("===== Chat Start =====") query = item["question"] - message_item = ScheduleMessageItem( - user_id=user_id, - mem_cube_id=mem_cube_id, - label=QUERY_LABEL, - mem_cube=mos.mem_cubes[mem_cube_id], - content=query, - timestamp=datetime.now(), - ) - mos.mem_scheduler.submit_messages(messages=message_item) - response = mos.chat(query, user_id=user_id) - message_item = ScheduleMessageItem( - user_id=user_id, - mem_cube_id=mem_cube_id, - label=ANSWER_LABEL, - mem_cube=mos.mem_cubes[mem_cube_id], - content=response, - timestamp=datetime.now(), - ) - mos.mem_scheduler.submit_messages(messages=message_item) - print(f"Query:\n {query}\n\nAnswer:\n {response}") + print(f"Query:\n {query}\n") + response = mos.chat(query=query, user_id=user_id) + print(f"Answer:\n {response}") + print("===== Chat End =====") show_web_logs(mem_scheduler=mos.mem_scheduler) @@ -236,6 +162,4 @@ def show_web_logs(mem_scheduler: GeneralScheduler): if __name__ == "__main__": - run_with_automatic_scheduler_init() - - run_with_manual_scheduler_init() + run_with_scheduler_init() diff --git a/examples/mem_user/user_manager_factory_example.py b/examples/mem_user/user_manager_factory_example.py new file mode 100644 index 000000000..ea50c30c9 --- /dev/null +++ b/examples/mem_user/user_manager_factory_example.py @@ -0,0 +1,111 @@ +"""Example demonstrating the use of UserManagerFactory with different backends.""" + +from memos.configs.mem_user import UserManagerConfigFactory +from memos.mem_user.factory import UserManagerFactory +from memos.mem_user.persistent_factory import PersistentUserManagerFactory + + +def example_sqlite_default(): + """Example: Create SQLite user manager with default settings.""" + print("=== SQLite Default Example ===") + + # Method 1: Using factory with minimal config + user_manager = UserManagerFactory.create_sqlite() + + # Method 2: Using config factory (equivalent) + UserManagerConfigFactory( + backend="sqlite", + config={}, # Uses all defaults + ) + + print(f"Created user manager: {type(user_manager).__name__}") + print(f"Database path: {user_manager.db_path}") + + # Test basic operations + users = user_manager.list_users() + print(f"Initial users: {[user.user_name for user in users]}") + + user_manager.close() + + +def example_sqlite_custom(): + """Example: Create SQLite user manager with custom settings.""" + print("\n=== SQLite Custom Example ===") + + config_factory = UserManagerConfigFactory( + backend="sqlite", config={"db_path": "/tmp/custom_memos.db", "user_id": "admin"} + ) + + user_manager = UserManagerFactory.from_config(config_factory) + print(f"Created user manager: {type(user_manager).__name__}") + print(f"Database path: {user_manager.db_path}") + + # Test operations + user_id = user_manager.create_user("test_user") + print(f"Created user: {user_id}") + + user_manager.close() + + +def example_mysql(): + """Example: Create MySQL user manager.""" + print("\n=== MySQL Example ===") + + # Method 1: Using factory with parameters + try: + user_manager = UserManagerFactory.create_mysql( + host="localhost", + port=3306, + username="root", + password="your_password", # Replace with actual password + database="test_memos_users", + ) + + print(f"Created user manager: {type(user_manager).__name__}") + print(f"Connection URL: {user_manager.connection_url}") + + # Test operations + users = user_manager.list_users() + print(f"Users: {[user.user_name for user in users]}") + + user_manager.close() + + except Exception as e: + print(f"MySQL connection failed (expected if not set up): {e}") + + +def example_persistent_managers(): + """Example: Create persistent user managers with configuration storage.""" + print("\n=== Persistent User Manager Examples ===") + + # SQLite persistent manager + config_factory = UserManagerConfigFactory(backend="sqlite", config={}) + + persistent_manager = PersistentUserManagerFactory.from_config(config_factory) + print(f"Created persistent manager: {type(persistent_manager).__name__}") + + # Test config operations + from memos.configs.mem_os import MOSConfig + + # Create a sample config (you might need to adjust this based on MOSConfig structure) + try: + # This is a simplified example - adjust based on actual MOSConfig requirements + sample_config = MOSConfig() # Use default config + + # Save user config + success = persistent_manager.save_user_config("test_user", sample_config) + print(f"Config saved: {success}") + + # Retrieve user config + retrieved_config = persistent_manager.get_user_config("test_user") + print(f"Config retrieved: {retrieved_config is not None}") + + except Exception as e: + print(f"Config operations failed: {e}") + + persistent_manager.close() + + +if __name__ == "__main__": + # Run all examples + example_sqlite_default() diff --git a/src/memos/api/config.py b/src/memos/api/config.py index 61e73f9ae..82dbf0d1d 100644 --- a/src/memos/api/config.py +++ b/src/memos/api/config.py @@ -115,6 +115,48 @@ def get_embedder_config() -> dict[str, Any]: }, } + @staticmethod + def get_internet_config() -> dict[str, Any]: + """Get embedder configuration.""" + return { + "backend": "xinyu", + "config": { + "api_key": os.getenv("XINYU_API_KEY"), + "search_engine_id": os.getenv("XINYU_SEARCH_ENGINE_ID"), + "max_results": 15, + "num_per_request": 10, + "reader": { + "backend": "simple_struct", + "config": { + "llm": { + "backend": "openai", + "config": { + "model_name_or_path": os.getenv("MEMRADER_MODEL"), + "temperature": 0.6, + "max_tokens": 5000, + "top_p": 0.95, + "top_k": 20, + "api_key": "EMPTY", + "api_base": os.getenv("MEMRADER_API_BASE"), + "remove_think_prefix": True, + "extra_body": {"chat_template_kwargs": {"enable_thinking": False}}, + }, + }, + "embedder": APIConfig.get_embedder_config(), + "chunker": { + "backend": "sentence", + "config": { + "tokenizer_or_token_counter": "gpt2", + "chunk_size": 512, + "chunk_overlap": 128, + "min_sentences_per_chunk": 1, + }, + }, + }, + }, + }, + } + @staticmethod def get_neo4j_community_config(user_id: str | None = None) -> dict[str, Any]: """Get Neo4j community configuration.""" @@ -212,6 +254,34 @@ def is_default_cube_config_enabled() -> bool: """Check if default cube config is enabled via environment variable.""" return os.getenv("MOS_ENABLE_DEFAULT_CUBE_CONFIG", "false").lower() == "true" + @staticmethod + def is_dingding_bot_enabled() -> bool: + """Check if DingDing bot is enabled via environment variable.""" + return os.getenv("ENABLE_DINGDING_BOT", "false").lower() == "true" + + @staticmethod + def get_dingding_bot_config() -> dict[str, Any] | None: + """Get DingDing bot configuration if enabled.""" + if not APIConfig.is_dingding_bot_enabled(): + return None + + return { + "enabled": True, + "access_token_user": os.getenv("DINGDING_ACCESS_TOKEN_USER", ""), + "secret_user": os.getenv("DINGDING_SECRET_USER", ""), + "access_token_error": os.getenv("DINGDING_ACCESS_TOKEN_ERROR", ""), + "secret_error": os.getenv("DINGDING_SECRET_ERROR", ""), + "robot_code": os.getenv("DINGDING_ROBOT_CODE", ""), + "app_key": os.getenv("DINGDING_APP_KEY", ""), + "app_secret": os.getenv("DINGDING_APP_SECRET", ""), + "oss_endpoint": os.getenv("OSS_ENDPOINT", ""), + "oss_region": os.getenv("OSS_REGION", ""), + "oss_bucket_name": os.getenv("OSS_BUCKET_NAME", ""), + "oss_access_key_id": os.getenv("OSS_ACCESS_KEY_ID", ""), + "oss_access_key_secret": os.getenv("OSS_ACCESS_KEY_SECRET", ""), + "oss_public_base_url": os.getenv("OSS_PUBLIC_BASE_URL", ""), + } + @staticmethod def get_product_default_config() -> dict[str, Any]: """Get default configuration for Product API.""" @@ -340,7 +410,6 @@ def create_user_config(user_name: str, user_id: str) -> tuple[MOSConfig, General "top_k": 30, "max_turns_window": 20, } - # Add scheduler configuration if enabled if APIConfig.is_scheduler_enabled(): config_dict["mem_scheduler"] = APIConfig.get_scheduler_config() @@ -352,7 +421,11 @@ def create_user_config(user_name: str, user_id: str) -> tuple[MOSConfig, General neo4j_community_config = APIConfig.get_neo4j_community_config(user_id) neo4j_config = APIConfig.get_neo4j_config(user_id) - + internet_config = ( + APIConfig.get_internet_config() + if os.getenv("ENABLE_INTERNET", "false").lower() == "true" + else None + ) graph_db_backend_map = { "neo4j-community": neo4j_community_config, "neo4j": neo4j_config, @@ -360,6 +433,7 @@ def create_user_config(user_name: str, user_id: str) -> tuple[MOSConfig, General graph_db_backend = os.getenv("NEO4J_BACKEND", "neo4j-community").lower() if graph_db_backend in graph_db_backend_map: # Create MemCube config + default_cube_config = GeneralMemCubeConfig.model_validate( { "user_id": user_id, @@ -374,6 +448,7 @@ def create_user_config(user_name: str, user_id: str) -> tuple[MOSConfig, General "config": graph_db_backend_map[graph_db_backend], }, "embedder": APIConfig.get_embedder_config(), + "internet_retriever": internet_config, }, }, "act_mem": {} @@ -384,7 +459,6 @@ def create_user_config(user_name: str, user_id: str) -> tuple[MOSConfig, General ) else: raise ValueError(f"Invalid Neo4j backend: {graph_db_backend}") - default_mem_cube = GeneralMemCube(default_cube_config) return default_config, default_mem_cube @@ -405,6 +479,11 @@ def get_default_cube_config() -> GeneralMemCubeConfig | None: "neo4j-community": neo4j_community_config, "neo4j": neo4j_config, } + internet_config = ( + APIConfig.get_internet_config() + if os.getenv("ENABLE_INTERNET", "false").lower() == "true" + else None + ) graph_db_backend = os.getenv("NEO4J_BACKEND", "neo4j-community").lower() if graph_db_backend in graph_db_backend_map: return GeneralMemCubeConfig.model_validate( @@ -423,6 +502,7 @@ def get_default_cube_config() -> GeneralMemCubeConfig | None: "embedder": APIConfig.get_embedder_config(), "reorganize": os.getenv("MOS_ENABLE_REORGANIZE", "false").lower() == "true", + "internet_retriever": internet_config, }, }, "act_mem": {} diff --git a/src/memos/api/context/context.py b/src/memos/api/context/context.py new file mode 100644 index 000000000..557ec84d1 --- /dev/null +++ b/src/memos/api/context/context.py @@ -0,0 +1,147 @@ +""" +Global request context management for trace_id and request-scoped data. + +This module provides optional trace_id functionality that can be enabled +when using the API components. It uses ContextVar to ensure thread safety +and request isolation. +""" + +import uuid + +from collections.abc import Callable +from contextvars import ContextVar +from typing import Any + + +# Global context variable for request-scoped data +_request_context: ContextVar[dict[str, Any] | None] = ContextVar("request_context", default=None) + + +class RequestContext: + """ + Request-scoped context object that holds trace_id and other request data. + + This provides a Flask g-like object for FastAPI applications. + """ + + def __init__(self, trace_id: str | None = None): + self.trace_id = trace_id or str(uuid.uuid4()) + self._data: dict[str, Any] = {} + + def set(self, key: str, value: Any) -> None: + """Set a value in the context.""" + self._data[key] = value + + def get(self, key: str, default: Any | None = None) -> Any: + """Get a value from the context.""" + return self._data.get(key, default) + + def __setattr__(self, name: str, value: Any) -> None: + if name.startswith("_") or name == "trace_id": + super().__setattr__(name, value) + else: + if not hasattr(self, "_data"): + super().__setattr__(name, value) + else: + self._data[name] = value + + def __getattr__(self, name: str) -> Any: + if hasattr(self, "_data") and name in self._data: + return self._data[name] + raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{name}'") + + def to_dict(self) -> dict[str, Any]: + """Convert context to dictionary.""" + return {"trace_id": self.trace_id, "data": self._data.copy()} + + +def set_request_context(context: RequestContext) -> None: + """ + Set the current request context. + + This is typically called by the API dependency injection system. + """ + _request_context.set(context.to_dict()) + + +def get_current_trace_id() -> str | None: + """ + Get the current request's trace_id. + + Returns: + The trace_id if available, None otherwise. + """ + context = _request_context.get() + if context: + return context.get("trace_id") + return None + + +def get_current_context() -> RequestContext | None: + """ + Get the current request context. + + Returns: + The current RequestContext if available, None otherwise. + """ + context_dict = _request_context.get() + if context_dict: + ctx = RequestContext(trace_id=context_dict.get("trace_id")) + ctx._data = context_dict.get("data", {}).copy() + return ctx + return None + + +def require_context() -> RequestContext: + """ + Get the current request context, raising an error if not available. + + Returns: + The current RequestContext. + + Raises: + RuntimeError: If called outside of a request context. + """ + context = get_current_context() + if context is None: + raise RuntimeError( + "No request context available. This function must be called within a request handler." + ) + return context + + +# Type for trace_id getter function +TraceIdGetter = Callable[[], str | None] + +# Global variable to hold the trace_id getter function +_trace_id_getter: TraceIdGetter | None = None + + +def set_trace_id_getter(getter: TraceIdGetter) -> None: + """ + Set a custom trace_id getter function. + + This allows the logging system to retrieve trace_id without importing + API-specific modules. + """ + global _trace_id_getter + _trace_id_getter = getter + + +def get_trace_id_for_logging() -> str | None: + """ + Get trace_id for logging purposes. + + This function is used by the logging system and will use either + the custom getter function or fall back to the default context. + """ + if _trace_id_getter: + try: + return _trace_id_getter() + except Exception: + pass + return get_current_trace_id() + + +# Initialize the default trace_id getter +set_trace_id_getter(get_current_trace_id) diff --git a/src/memos/api/context/dependencies.py b/src/memos/api/context/dependencies.py new file mode 100644 index 000000000..d26cadaa5 --- /dev/null +++ b/src/memos/api/context/dependencies.py @@ -0,0 +1,90 @@ +import logging + +from fastapi import Depends, Header, Request + +from memos.api.context.context import RequestContext, set_request_context + + +logger = logging.getLogger(__name__) + +# Type alias for the RequestContext from context module +G = RequestContext + + +def get_trace_id_from_header( + trace_id: str | None = Header(None, alias="trace-id"), + x_trace_id: str | None = Header(None, alias="x-trace-id"), + g_trace_id: str | None = Header(None, alias="g-trace-id"), +) -> str | None: + """ + Extract trace_id from various possible headers. + + Priority: g-trace-id > x-trace-id > trace-id + """ + return g_trace_id or x_trace_id or trace_id + + +def get_request_context( + request: Request, trace_id: str | None = Depends(get_trace_id_from_header) +) -> RequestContext: + """ + Get request context object with trace_id and request metadata. + + This function creates a RequestContext and automatically sets it + in the global context for use throughout the request lifecycle. + """ + # Create context object + ctx = RequestContext(trace_id=trace_id) + + # Set the context globally for this request + set_request_context(ctx) + + # Log request start + logger.info(f"Request started with trace_id: {ctx.trace_id}") + + # Add request metadata to context + ctx.set("method", request.method) + ctx.set("path", request.url.path) + ctx.set("client_ip", request.client.host if request.client else None) + + return ctx + + +def get_g_object(trace_id: str | None = Depends(get_trace_id_from_header)) -> G: + """ + Get Flask g-like object for the current request. + + This creates a RequestContext and sets it globally for access + throughout the request lifecycle. + """ + g = RequestContext(trace_id=trace_id) + set_request_context(g) + logger.info(f"Request g object created with trace_id: {g.trace_id}") + return g + + +def get_current_g() -> G | None: + """ + Get the current request's g object from anywhere in the application. + + Returns: + The current request's g object if available, None otherwise. + """ + from memos.context import get_current_context + + return get_current_context() + + +def require_g() -> G: + """ + Get the current request's g object, raising an error if not available. + + Returns: + The current request's g object. + + Raises: + RuntimeError: If called outside of a request context. + """ + from memos.context import require_context + + return require_context() diff --git a/src/memos/api/product_models.py b/src/memos/api/product_models.py index 0e9d5ff59..f21ba2b6b 100644 --- a/src/memos/api/product_models.py +++ b/src/memos/api/product_models.py @@ -83,6 +83,7 @@ class ChatRequest(BaseRequest): query: str = Field(..., description="Chat query message") mem_cube_id: str | None = Field(None, description="Cube ID to use for chat") history: list[MessageDict] | None = Field(None, description="Chat history") + internet_search: bool = Field(True, description="Whether to use internet search") class UserCreate(BaseRequest): @@ -150,6 +151,7 @@ class SearchRequest(BaseRequest): user_id: str = Field(..., description="User ID") query: str = Field(..., description="Search query") mem_cube_id: str | None = Field(None, description="Cube ID to search in") + top_k: int = Field(10, description="Number of results to return") class SuggestionRequest(BaseRequest): diff --git a/src/memos/api/routers/product_router.py b/src/memos/api/routers/product_router.py index 92acb38a4..48ecb10f5 100644 --- a/src/memos/api/routers/product_router.py +++ b/src/memos/api/routers/product_router.py @@ -2,10 +2,14 @@ import logging import traceback -from fastapi import APIRouter, HTTPException +from datetime import datetime +from typing import Annotated + +from fastapi import APIRouter, Depends, HTTPException from fastapi.responses import StreamingResponse from memos.api.config import APIConfig +from memos.api.context.dependencies import G, get_g_object from memos.api.product_models import ( BaseResponse, ChatRequest, @@ -22,6 +26,7 @@ ) from memos.configs.mem_os import MOSConfig from memos.mem_os.product import MOSProduct +from memos.memos_tools.notification_service import get_error_bot_function, get_online_bot_function logger = logging.getLogger(__name__) @@ -37,16 +42,25 @@ def get_mos_product_instance(): global MOS_PRODUCT_INSTANCE if MOS_PRODUCT_INSTANCE is None: default_config = APIConfig.get_product_default_config() - print(default_config) + logger.info(f"*********init_default_mos_config********* {default_config}") from memos.configs.mem_os import MOSConfig mos_config = MOSConfig(**default_config) # Get default cube config from APIConfig (may be None if disabled) default_cube_config = APIConfig.get_default_cube_config() - print("*********default_cube_config*********", default_cube_config) + logger.info(f"*********initdefault_cube_config******** {default_cube_config}") + + # Get DingDing bot functions + dingding_enabled = APIConfig.is_dingding_bot_enabled() + online_bot = get_online_bot_function() if dingding_enabled else None + error_bot = get_error_bot_function() if dingding_enabled else None + MOS_PRODUCT_INSTANCE = MOSProduct( - default_config=mos_config, default_cube_config=default_cube_config + default_config=mos_config, + default_cube_config=default_cube_config, + online_bot=online_bot, + error_bot=error_bot, ) logger.info("MOSProduct instance created successfully with inheritance architecture") return MOS_PRODUCT_INSTANCE @@ -56,7 +70,7 @@ def get_mos_product_instance(): @router.post("/configure", summary="Configure MOSProduct", response_model=SimpleResponse) -async def set_config(config): +def set_config(config): """Set MOSProduct configuration.""" global MOS_PRODUCT_INSTANCE MOS_PRODUCT_INSTANCE = MOSProduct(default_config=config) @@ -64,9 +78,18 @@ async def set_config(config): @router.post("/users/register", summary="Register a new user", response_model=UserRegisterResponse) -async def register_user(user_req: UserRegisterRequest): +def register_user(user_req: UserRegisterRequest, g: Annotated[G, Depends(get_g_object)]): """Register a new user with configuration and default cube.""" try: + # Set request-related information in g object + g.user_id = user_req.user_id + g.action = "user_register" + g.timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + + logger.info(f"Starting user registration for user_id: {user_req.user_id}") + logger.info(f"Request trace_id: {g.trace_id}") + logger.info(f"Request timestamp: {g.timestamp}") + # Get configuration for the user user_config, default_mem_cube = APIConfig.create_user_config( user_name=user_req.user_id, user_id=user_req.user_id @@ -100,7 +123,7 @@ async def register_user(user_req: UserRegisterRequest): @router.get( "/suggestions/{user_id}", summary="Get suggestion queries", response_model=SuggestionResponse ) -async def get_suggestion_queries(user_id: str): +def get_suggestion_queries(user_id: str): """Get suggestion queries for a specific user.""" try: mos_product = get_mos_product_instance() @@ -120,7 +143,7 @@ async def get_suggestion_queries(user_id: str): summary="Get suggestion queries with language", response_model=SuggestionResponse, ) -async def get_suggestion_queries_post(suggestion_req: SuggestionRequest): +def get_suggestion_queries_post(suggestion_req: SuggestionRequest): """Get suggestion queries for a specific user with language preference.""" try: mos_product = get_mos_product_instance() @@ -138,7 +161,7 @@ async def get_suggestion_queries_post(suggestion_req: SuggestionRequest): @router.post("/get_all", summary="Get all memories for user", response_model=MemoryResponse) -async def get_all_memories(memory_req: GetMemoryRequest): +def get_all_memories(memory_req: GetMemoryRequest): """Get all memories for a specific user.""" try: mos_product = get_mos_product_instance() @@ -165,7 +188,7 @@ async def get_all_memories(memory_req: GetMemoryRequest): @router.post("/add", summary="add a new memory", response_model=SimpleResponse) -async def create_memory(memory_req: MemoryCreateRequest): +def create_memory(memory_req: MemoryCreateRequest): """Create a new memory for a specific user.""" try: mos_product = get_mos_product_instance() @@ -186,7 +209,7 @@ async def create_memory(memory_req: MemoryCreateRequest): @router.post("/search", summary="Search memories", response_model=SearchResponse) -async def search_memories(search_req: SearchRequest): +def search_memories(search_req: SearchRequest): """Search memories for a specific user.""" try: mos_product = get_mos_product_instance() @@ -194,6 +217,7 @@ async def search_memories(search_req: SearchRequest): query=search_req.query, user_id=search_req.user_id, install_cube_ids=[search_req.mem_cube_id] if search_req.mem_cube_id else None, + top_k=search_req.top_k, ) return SearchResponse(message="Search completed successfully", data=result) @@ -205,24 +229,23 @@ async def search_memories(search_req: SearchRequest): @router.post("/chat", summary="Chat with MemOS") -async def chat(chat_req: ChatRequest): +def chat(chat_req: ChatRequest): """Chat with MemOS for a specific user. Returns SSE stream.""" try: mos_product = get_mos_product_instance() - async def generate_chat_response(): + def generate_chat_response(): """Generate chat response as SSE stream.""" try: - import asyncio - - for chunk in mos_product.chat_with_references( + # Directly yield from the generator without async wrapper + yield from mos_product.chat_with_references( query=chat_req.query, user_id=chat_req.user_id, cube_id=chat_req.mem_cube_id, history=chat_req.history, - ): - yield chunk - await asyncio.sleep(0.00001) # 50ms delay between chunks + internet_search=chat_req.internet_search, + ) + except Exception as e: logger.error(f"Error in chat stream: {e}") error_data = f"data: {json.dumps({'type': 'error', 'content': str(traceback.format_exc())})}\n\n" @@ -230,11 +253,14 @@ async def generate_chat_response(): return StreamingResponse( generate_chat_response(), - media_type="text/plain", + media_type="text/event-stream", headers={ "Cache-Control": "no-cache", "Connection": "keep-alive", "Content-Type": "text/event-stream", + "Access-Control-Allow-Origin": "*", + "Access-Control-Allow-Headers": "*", + "Access-Control-Allow-Methods": "*", }, ) @@ -246,7 +272,7 @@ async def generate_chat_response(): @router.get("/users", summary="List all users", response_model=BaseResponse[list]) -async def list_users(): +def list_users(): """List all registered users.""" try: mos_product = get_mos_product_instance() @@ -274,7 +300,7 @@ async def get_user_info(user_id: str): @router.get( "/configure/{user_id}", summary="Get MOSProduct configuration", response_model=SimpleResponse ) -async def get_config(user_id: str): +def get_config(user_id: str): """Get MOSProduct configuration.""" global MOS_PRODUCT_INSTANCE config = MOS_PRODUCT_INSTANCE.default_config @@ -284,7 +310,7 @@ async def get_config(user_id: str): @router.get( "/users/{user_id}/config", summary="Get user configuration", response_model=BaseResponse[dict] ) -async def get_user_config(user_id: str): +def get_user_config(user_id: str): """Get user-specific configuration.""" try: mos_product = get_mos_product_instance() @@ -308,7 +334,7 @@ async def get_user_config(user_id: str): @router.put( "/users/{user_id}/config", summary="Update user configuration", response_model=SimpleResponse ) -async def update_user_config(user_id: str, config_data: dict): +def update_user_config(user_id: str, config_data: dict): """Update user-specific configuration.""" try: mos_product = get_mos_product_instance() @@ -333,7 +359,7 @@ async def update_user_config(user_id: str, config_data: dict): @router.get( "/instances/status", summary="Get user configuration status", response_model=BaseResponse[dict] ) -async def get_instance_status(): +def get_instance_status(): """Get information about active user configurations in memory.""" try: mos_product = get_mos_product_instance() @@ -347,7 +373,7 @@ async def get_instance_status(): @router.get("/instances/count", summary="Get active user count", response_model=BaseResponse[int]) -async def get_active_user_count(): +def get_active_user_count(): """Get the number of active user configurations in memory.""" try: mos_product = get_mos_product_instance() diff --git a/src/memos/configs/graph_db.py b/src/memos/configs/graph_db.py index 880cda87c..01246b06e 100644 --- a/src/memos/configs/graph_db.py +++ b/src/memos/configs/graph_db.py @@ -9,7 +9,7 @@ class BaseGraphDBConfig(BaseConfig): """Base class for all graph database configurations.""" - uri: str + uri: str | list user: str password: str @@ -103,7 +103,7 @@ def validate_community(self): return self -class NebulaGraphDBConfig(BaseConfig): +class NebulaGraphDBConfig(BaseGraphDBConfig): """ NebulaGraph-specific configuration. @@ -121,8 +121,6 @@ class NebulaGraphDBConfig(BaseConfig): user_name = "alice" """ - password: str - hosts: list[str] = Field(..., description="List of host:port strings for NebulaGraph servers") space: str = Field( ..., description="The name of the target NebulaGraph space (like a database)" ) diff --git a/src/memos/configs/internet_retriever.py b/src/memos/configs/internet_retriever.py index 56f892ac9..96b731c75 100644 --- a/src/memos/configs/internet_retriever.py +++ b/src/memos/configs/internet_retriever.py @@ -6,6 +6,7 @@ from memos.configs.base import BaseConfig from memos.exceptions import ConfigurationError +from memos.mem_reader.factory import MemReaderConfigFactory class BaseInternetRetrieverConfig(BaseConfig): @@ -47,6 +48,11 @@ class XinyuSearchConfig(BaseInternetRetrieverConfig): num_per_request: int = Field( default=10, description="Number of results per API request (not used for Xinyu)" ) + reader: MemReaderConfigFactory = Field( + ..., + default_factory=MemReaderConfigFactory, + description="Reader configuration", + ) class InternetRetrieverConfigFactory(BaseConfig): diff --git a/src/memos/configs/mem_os.py b/src/memos/configs/mem_os.py index 96b4094ed..0645fce44 100644 --- a/src/memos/configs/mem_os.py +++ b/src/memos/configs/mem_os.py @@ -8,6 +8,7 @@ from memos.configs.llm import LLMConfigFactory from memos.configs.mem_reader import MemReaderConfigFactory from memos.configs.mem_scheduler import SchedulerConfigFactory +from memos.configs.mem_user import UserManagerConfigFactory class MOSConfig(BaseConfig): @@ -33,6 +34,10 @@ class MOSConfig(BaseConfig): default=None, description="Memory scheduler configuration for managing memory operations", ) + user_manager: UserManagerConfigFactory = Field( + default_factory=lambda: UserManagerConfigFactory(backend="sqlite", config={}), + description="User manager configuration for database operations", + ) max_turns_window: int = Field( default=15, description="Maximum number of turns to keep in the conversation history", diff --git a/src/memos/configs/mem_reader.py b/src/memos/configs/mem_reader.py index c867f30b5..1c62087a3 100644 --- a/src/memos/configs/mem_reader.py +++ b/src/memos/configs/mem_reader.py @@ -15,6 +15,15 @@ class BaseMemReaderConfig(BaseConfig): created_at: datetime = Field( default_factory=datetime.now, description="Creation timestamp for the MemReader" ) + + @field_validator("created_at", mode="before") + @classmethod + def parse_datetime(cls, value): + """Parse datetime from string if needed.""" + if isinstance(value, str): + return datetime.fromisoformat(value.replace("Z", "+00:00")) + return value + llm: LLMConfigFactory = Field(..., description="LLM configuration for the MemReader") embedder: EmbedderConfigFactory = Field( ..., description="Embedder configuration for the MemReader" diff --git a/src/memos/configs/mem_user.py b/src/memos/configs/mem_user.py new file mode 100644 index 000000000..3ff1066e5 --- /dev/null +++ b/src/memos/configs/mem_user.py @@ -0,0 +1,58 @@ +from typing import Any, ClassVar + +from pydantic import BaseModel, Field, field_validator, model_validator + +from memos.configs.base import BaseConfig + + +class BaseUserManagerConfig(BaseConfig): + """Base configuration class for user managers.""" + + user_id: str = Field(default="root", description="Default user ID for initialization") + + +class SQLiteUserManagerConfig(BaseUserManagerConfig): + """SQLite user manager configuration.""" + + db_path: str | None = Field( + default=None, + description="Path to SQLite database file. If None, uses default path in MEMOS_DIR", + ) + + +class MySQLUserManagerConfig(BaseUserManagerConfig): + """MySQL user manager configuration.""" + + host: str = Field(default="localhost", description="MySQL server host") + port: int = Field(default=3306, description="MySQL server port") + username: str = Field(default="root", description="MySQL username") + password: str = Field(default="", description="MySQL password") + database: str = Field(default="memos_users", description="MySQL database name") + charset: str = Field(default="utf8mb4", description="MySQL charset") + + +class UserManagerConfigFactory(BaseModel): + """Factory for user manager configurations.""" + + backend: str = Field(default="sqlite", description="Backend for user manager") + config: dict[str, Any] = Field( + default_factory=dict, description="Configuration for the user manager backend" + ) + + backend_to_class: ClassVar[dict[str, Any]] = { + "sqlite": SQLiteUserManagerConfig, + "mysql": MySQLUserManagerConfig, + } + + @field_validator("backend") + @classmethod + def validate_backend(cls, backend: str) -> str: + if backend not in cls.backend_to_class: + raise ValueError(f"Unsupported user manager backend: {backend}") + return backend + + @model_validator(mode="after") + def instantiate_config(self): + config_class = self.backend_to_class[self.backend] + self.config = config_class(**self.config) + return self diff --git a/src/memos/graph_dbs/base.py b/src/memos/graph_dbs/base.py index d59139ef2..e8b331395 100644 --- a/src/memos/graph_dbs/base.py +++ b/src/memos/graph_dbs/base.py @@ -146,7 +146,7 @@ def search_by_embedding(self, vector: list[float], top_k: int = 5) -> list[dict] """ @abstractmethod - def get_by_metadata(self, filters: dict[str, Any]) -> list[str]: + def get_by_metadata(self, filters: list[dict[str, Any]]) -> list[str]: """ Retrieve node IDs that match given metadata filters. @@ -162,6 +162,14 @@ def get_by_metadata(self, filters: dict[str, Any]) -> list[str]: - Can be used for faceted recall or prefiltering before embedding rerank. """ + @abstractmethod + def get_structure_optimization_candidates(self, scope: str) -> list[dict]: + """ + Find nodes that are likely candidates for structure optimization: + - Isolated nodes, nodes with empty background, or nodes with exactly one child. + - Plus: the child of any parent node that has exactly one child. + """ + # Structure Maintenance @abstractmethod def deduplicate_nodes(self) -> None: diff --git a/src/memos/graph_dbs/nebular.py b/src/memos/graph_dbs/nebular.py index 24647c0bc..9fca1988e 100644 --- a/src/memos/graph_dbs/nebular.py +++ b/src/memos/graph_dbs/nebular.py @@ -1,8 +1,12 @@ +import traceback + +from contextlib import suppress from datetime import datetime +from queue import Empty, Queue +from threading import Lock from typing import Any, Literal -from nebulagraph_python.py_data_types import NVector -from nebulagraph_python.value_wrapper import ValueWrapper +import numpy as np from memos.configs.graph_db import NebulaGraphDBConfig from memos.dependency import require_python_package @@ -13,6 +17,12 @@ logger = get_logger(__name__) +def _normalize(vec: list[float]) -> list[float]: + v = np.asarray(vec, dtype=np.float32) + norm = np.linalg.norm(v) + return (v / (norm if norm else 1.0)).tolist() + + def _compose_node(item: dict[str, Any]) -> tuple[str, str, dict[str, Any]]: node_id = item["id"] memory = item["memory"] @@ -36,7 +46,7 @@ def _prepare_node_metadata(metadata: dict[str, Any]) -> dict[str, Any]: # Normalize embedding type embedding = metadata.get("embedding") if embedding and isinstance(embedding, list): - metadata["embedding"] = [float(x) for x in embedding] + metadata["embedding"] = _normalize([float(x) for x in embedding]) return metadata @@ -46,6 +56,8 @@ def _escape_str(value: str) -> str: def _format_value(val: Any, key: str = "") -> str: + from nebulagraph_python.py_data_types import NVector + if isinstance(val, str): return f'"{_escape_str(val)}"' elif isinstance(val, (int | float)): @@ -77,6 +89,86 @@ def _format_datetime(value: str | datetime) -> str: return str(value) +class SessionPoolError(Exception): + pass + + +class SessionPool: + @require_python_package( + import_name="nebulagraph_python", + install_command="pip install ... @Tianxing", + install_link=".....", + ) + def __init__( + self, + hosts: list[str], + user: str, + password: str, + minsize: int = 1, + maxsize: int = 10000, + ): + self.hosts = hosts + self.user = user + self.password = password + self.maxsize = maxsize + self.pool = Queue(maxsize) + self.lock = Lock() + + self.clients = [] + + for _ in range(minsize): + self._create_and_add_client() + + def _create_and_add_client(self): + from nebulagraph_python import NebulaClient + + client = NebulaClient(self.hosts, self.user, self.password) + self.pool.put(client) + self.clients.append(client) + + def get_client(self, timeout: float = 5.0): + from nebulagraph_python import NebulaClient + + try: + return self.pool.get(timeout=timeout) + except Empty: + with self.lock: + if len(self.clients) < self.maxsize: + client = NebulaClient(self.hosts, self.user, self.password) + self.clients.append(client) + return client + raise RuntimeError("NebulaClientPool exhausted") from None + + def return_client(self, client): + self.pool.put(client) + + def close(self): + for client in self.clients: + with suppress(Exception): + client.close() + self.clients.clear() + + def get(self): + """ + Context manager: with pool.get() as client: + """ + + class _ClientContext: + def __init__(self, outer): + self.outer = outer + self.client = None + + def __enter__(self): + self.client = self.outer.get_client() + return self.client + + def __exit__(self, exc_type, exc_val, exc_tb): + if self.client: + self.outer.return_client(self.client) + + return _ClientContext(self) + + class NebulaGraphDB(BaseGraphDB): """ NebulaGraph-based implementation of a graph memory store. @@ -95,7 +187,7 @@ def __init__(self, config: NebulaGraphDBConfig): - hosts: list[str] like ["host1:port", "host2:port"] - user: str - password: str - - space: str (optional for basic commands) + - db_name: str (optional for basic commands) Example config: { @@ -105,26 +197,38 @@ def __init__(self, config: NebulaGraphDBConfig): "space": "test" } """ - from nebulagraph_python.client import NebulaClient self.config = config - self.client = NebulaClient( - hosts=config.get("hosts"), - username=config.get("user_name"), - password=config.get("password"), - ) self.db_name = config.space - self.space = config.get("space") self.user_name = config.user_name self.system_db_name = "system" if config.use_multi_db else config.space + self.pool = SessionPool( + hosts=config.get("uri"), + user=config.get("user"), + password=config.get("password"), + minsize=1, + maxsize=config.get("max_client", 1000), + ) + if config.auto_create: self._ensure_database_exists() + self.execute_query(f"SESSION SET GRAPH `{self.db_name}`") + # Create only if not exists self.create_index(dimensions=config.embedding_dimension) logger.info("Connected to NebulaGraph successfully.") + def execute_query(self, gql: str, timeout: float = 5.0, auto_set_db: bool = True): + with self.pool.get() as client: + if auto_set_db and self.db_name: + client.execute(f"SESSION SET GRAPH `{self.db_name}`") + return client.execute(gql, timeout=timeout) + + def close(self): + self.pool.close() + def create_index( self, label: str = "Memory", @@ -132,19 +236,10 @@ def create_index( dimensions: int = 3072, index_name: str = "memory_vector_index", ) -> None: - create_vector_index = f""" - CREATE VECTOR INDEX IF NOT EXISTS {index_name} - ON NODE Memory::{vector_property} - OPTIONS {{ - DIM: {dimensions}, - METRIC: L2, - TYPE: IVF, - NLIST: 100, - TRAINSIZE: 1000 - }} - FOR memory_graph - """ - self.client.execute(create_vector_index) + # Create vector index + self._create_vector_index(label, vector_property, dimensions, index_name) + # Create indexes + self._create_basic_property_indexes() def remove_oldest_memory(self, memory_type: str, keep_latest: int) -> None: """ @@ -166,7 +261,7 @@ def remove_oldest_memory(self, memory_type: str, keep_latest: int) -> None: OFFSET {keep_latest} DETACH DELETE n """ - self.client.execute(query) + self.execute_query(query) def add_node(self, id: str, memory: str, metadata: dict[str, Any]) -> None: """ @@ -183,13 +278,39 @@ def add_node(self, id: str, memory: str, metadata: dict[str, Any]) -> None: metadata["id"] = id metadata["memory"] = memory + if "embedding" in metadata and isinstance(metadata["embedding"], list): + metadata["embedding"] = _normalize(metadata["embedding"]) + properties = ", ".join(f"{k}: {_format_value(v, k)}" for k, v in metadata.items()) gql = f"INSERT OR IGNORE (n@Memory {{{properties}}})" try: - self.client.execute(gql) + self.execute_query(gql) + logger.info("insert success") + except Exception as e: + logger.error( + f"Failed to insert vertex {id}: gql: {gql}, {e}\ntrace: {traceback.format_exc()}" + ) + + def node_not_exist(self, scope: str) -> int: + if not self.config.use_multi_db and self.config.user_name: + filter_clause = f'n.memory_type = "{scope}" AND n.user_name = "{self.config.user_name}"' + else: + filter_clause = f'n.memory_type = "{scope}"' + + query = f""" + MATCH (n@Memory) + WHERE {filter_clause} + RETURN n + LIMIT 1 + """ + + try: + result = self.execute_query(query) + return result.size == 0 except Exception as e: - logger.error(f"Failed to insert vertex {id}: {e}") + logger.error(f"[node_not_exist] Query failed: {e}", exc_info=True) + raise def update_node(self, id: str, fields: dict[str, Any]) -> None: """ @@ -210,7 +331,7 @@ def update_node(self, id: str, fields: dict[str, Any]) -> None: query += f'WHERE n.user_name = "{self.config.user_name}"' query += f"\nSET {set_clause_str}" - self.client.execute(query) + self.execute_query(query) def delete_node(self, id: str) -> None: """ @@ -225,7 +346,7 @@ def delete_node(self, id: str) -> None: user_name = self.config.user_name query += f" WHERE n.user_name = {_format_value(user_name)}" query += "\n DETACH DELETE n" - self.client.execute(query) + self.execute_query(query) def add_edge(self, source_id: str, target_id: str, type: str): """ @@ -247,7 +368,7 @@ def add_edge(self, source_id: str, target_id: str, type: str): INSERT (a) -[e@{type} {props}]-> (b) ''' try: - self.client.execute(insert_stmt) + self.execute_query(insert_stmt) except Exception as e: logger.error(f"Failed to insert edge: {e}", exc_info=True) @@ -269,12 +390,12 @@ def delete_edge(self, source_id: str, target_id: str, type: str) -> None: query += f" AND a.user_name = {_format_value(user_name)} AND b.user_name = {_format_value(user_name)}" query += "\nDELETE r" - self.client.execute(query) + self.execute_query(query) def get_memory_count(self, memory_type: str) -> int: query = f""" MATCH (n@Memory) - WHERE n.memory_type = {memory_type} + WHERE n.memory_type = "{memory_type}" """ if not self.config.use_multi_db and self.config.user_name: user_name = self.config.user_name @@ -282,7 +403,7 @@ def get_memory_count(self, memory_type: str) -> int: query += "\nRETURN COUNT(n) AS count" try: - result = self.client.execute(query) + result = self.execute_query(query) return result.one_or_none()["count"].value except Exception as e: logger.error(f"[get_memory_count] Failed: {e}") @@ -291,14 +412,14 @@ def get_memory_count(self, memory_type: str) -> int: def count_nodes(self, scope: str) -> int: query = f""" MATCH (n@Memory) - WHERE n.memory_type = {scope} + WHERE n.memory_type = "{scope}" """ if not self.config.use_multi_db and self.config.user_name: user_name = self.config.user_name query += f"\nAND n.user_name = '{user_name}'" query += "\nRETURN count(n) AS count" - result = self.client.execute(query) + result = self.execute_query(query) return result.one_or_none()["count"].value def edge_exists( @@ -336,8 +457,11 @@ def edge_exists( query += "\nRETURN r" # Run the Cypher query - result = self.client.execute(query) - return result.one_or_none().values() is not None + result = self.execute_query(query) + record = result.one_or_none() + if record is None: + return False + return record.values() is not None # Graph Query & Reasoning def get_node(self, id: str) -> dict[str, Any] | None: @@ -351,21 +475,21 @@ def get_node(self, id: str) -> dict[str, Any] | None: dict: Node properties as key-value pairs, or None if not found. """ gql = f""" - USE memory_graph + USE `{self.db_name}` MATCH (v {{id: '{id}'}}) RETURN v """ try: - result = self.client.execute(gql) + result = self.execute_query(gql) record = result.one_or_none() if record is None: return None node_wrapper = record["v"].as_node() props = node_wrapper.get_properties() - - return {key: self._parse_node(val) for key, val in props.items()} + node = self._parse_node(props) + return node except Exception as e: logger.error(f"[get_node] Failed to retrieve node '{id}': {e}") @@ -392,8 +516,13 @@ def get_nodes(self, ids: list[str]) -> list[dict[str, Any]]: query = f"MATCH (n@Memory) WHERE n.id IN {ids} {where_user} RETURN n" - results = self.client.execute(query) - return [self._parse_node(record["n"]) for record in results] + results = self.execute_query(query) + nodes = [] + for rec in results: + node_props = rec["n"].as_node().get_properties() + nodes.append(self._parse_node(node_props)) + + return nodes def get_edges(self, id: str, type: str = "ANY", direction: str = "ANY") -> list[dict[str, str]]: """ @@ -423,7 +552,7 @@ def get_edges(self, id: str, type: str = "ANY", direction: str = "ANY") -> list[ where_clause = f"a.id = '{id}'" elif direction == "ANY": pattern = f"(a@Memory)-[r{rel_type}]-(b@Memory)" - where_clause = f"a.id = {id} OR b.id = {id}" + where_clause = f"a.id = '{id}' OR b.id = '{id}'" else: raise ValueError("Invalid direction. Must be 'OUTGOING', 'INCOMING', or 'ANY'.") @@ -431,12 +560,12 @@ def get_edges(self, id: str, type: str = "ANY", direction: str = "ANY") -> list[ where_clause += f" AND a.user_name = '{self.config.user_name}' AND b.user_name = '{self.config.user_name}'" query = f""" - MATCH {pattern} - WHERE {where_clause} - RETURN a.id AS from_id, b.id AS to_id, type(r) AS edge_type - """ + MATCH {pattern} + WHERE {where_clause} + RETURN a.id AS from_id, b.id AS to_id, type(r) AS edge_type + """ - result = self.client.execute(query) + result = self.execute_query(query) edges = [] for record in result: edges.append( @@ -448,7 +577,6 @@ def get_edges(self, id: str, type: str = "ANY", direction: str = "ANY") -> list[ ) return edges - # TODO def get_neighbors_by_tag( self, tags: list[str], @@ -468,27 +596,50 @@ def get_neighbors_by_tag( Returns: List of dicts with node details and overlap count. """ - where_user = "" + if not tags: + return [] + + where_clauses = [ + 'n.status = "activated"', + 'NOT (n.node_type = "reasoning")', + 'NOT (n.memory_type = "WorkingMemory")', + ] + if exclude_ids: + where_clauses.append(f"NOT (n.id IN {exclude_ids})") + if not self.config.use_multi_db and self.config.user_name: - user_name = self.config.user_name - where_user = f"AND n.user_name = {user_name}" + where_clauses.append(f'n.user_name = "{self.config.user_name}"') + + where_clause = " AND ".join(where_clauses) + tag_list_literal = "[" + ", ".join(f'"{_escape_str(t)}"' for t in tags) + "]" query = f""" - MATCH (n@Memory) - LET overlap_tags = [tag IN n.tags WHERE tag IN {tags}] - WHERE NOT n.id IN {exclude_ids} - AND n.status = 'activated' - AND n.node_type <> 'reasoning' - AND n.memory_type <> 'WorkingMemory' - {where_user} - AND size(overlap_tags) >= {min_overlap} - RETURN n, size(overlap_tags) AS overlap_count - ORDER BY overlap_count DESC - LIMIT {top_k} - """ - print(query) - result = self.client.execute(query) - return [self._parse_node(dict(record)) for record in result] + LET tag_list = {tag_list_literal} + + MATCH (n@Memory) + WHERE {where_clause} + RETURN n, + size( filter( n.tags, t -> t IN tag_list ) ) AS overlap_count + ORDER BY overlap_count DESC + LIMIT {top_k} + """ + + result = self.execute_query(query) + neighbors: list[dict[str, Any]] = [] + for r in result: + node_props = r["n"].as_node().get_properties() + parsed = self._parse_node(node_props) # --> {id, memory, metadata} + + parsed["overlap_count"] = r["overlap_count"].value + neighbors.append(parsed) + + neighbors.sort(key=lambda x: x["overlap_count"], reverse=True) + neighbors = neighbors[:top_k] + result = [] + for neighbor in neighbors[:top_k]: + neighbor.pop("overlap_count") + result.append(neighbor) + return result def get_children_with_embeddings(self, id: str) -> list[dict[str, Any]]: where_user = "" @@ -498,15 +649,20 @@ def get_children_with_embeddings(self, id: str) -> list[dict[str, Any]]: where_user = f"AND p.user_name = '{user_name}' AND c.user_name = '{user_name}'" query = f""" - MATCH (p@Memory)-[@PARENT]->(c@Memory) - WHERE p.id = "{id}" {where_user} - RETURN c.id AS id, c.embedding AS embedding, c.memory AS memory - """ - result = self.client.execute(query) - return [ - {"id": r["id"].value, "embedding": r["embedding"].value, "memory": r["memory"].value} - for r in result - ] + MATCH (p@Memory)-[@PARENT]->(c@Memory) + WHERE p.id = "{id}" {where_user} + RETURN c.id AS id, c.embedding AS embedding, c.memory AS memory + """ + result = self.execute_query(query) + children = [] + for row in result: + eid = row["id"].value # STRING + emb_v = row["embedding"].value # NVector + emb = list(emb_v.values) if emb_v else [] + mem = row["memory"].value # STRING + + children.append({"id": eid, "embedding": emb, "memory": mem}) + return children def get_subgraph( self, center_id: str, depth: int = 2, center_status: str = "activated" @@ -540,21 +696,30 @@ def get_subgraph( collect(EDGES(p)) AS edge_chains """ - result = self.client.execute(gql).one_or_none() + result = self.execute_query(gql).one_or_none() if not result or result.size == 0: return {"core_node": None, "neighbors": [], "edges": []} - core_node = self._parse_node(result["center"]) - neighbors = [self._parse_node(n) for n in result["neighbors"].value] + core_node_props = result["center"].as_node().get_properties() + core_node = self._parse_node(core_node_props) + neighbors = [] + vid_to_id_map = {result["center"].as_node().node_id: core_node["id"]} + for n in result["neighbors"].value: + n_node = n.as_node() + n_props = n_node.get_properties() + node_parsed = self._parse_node(n_props) + neighbors.append(node_parsed) + vid_to_id_map[n_node.node_id] = node_parsed["id"] + edges = [] - for rel_chains in result["edge_chains"].value: - for chain in rel_chains.value: - edge = chain.value + for chain_group in result["edge_chains"].value: + for edge_wr in chain_group.value: + edge = edge_wr.value edges.append( { "type": edge.get_type(), - "source": edge.get_src_id(), - "target": edge.get_dst_id(), + "source": vid_to_id_map.get(edge.get_src_id()), + "target": vid_to_id_map.get(edge.get_dst_id()), } ) @@ -591,6 +756,7 @@ def search_by_embedding( - Typical use case: restrict to 'status = activated' to avoid matching archived or merged nodes. """ + vector = _normalize(vector) dim = len(vector) vector_str = ",".join(f"{float(x)}" for x in vector) gql_vector = f"VECTOR<{dim}, FLOAT>([{vector_str}])" @@ -606,18 +772,18 @@ def search_by_embedding( where_clause = f"WHERE {' AND '.join(where_clauses)}" if where_clauses else "" gql = f""" - USE memory_graph + USE `{self.db_name}` MATCH (n@Memory) {where_clause} - ORDER BY euclidean(n.embedding, {gql_vector}) ASC + ORDER BY inner_product(n.embedding, {gql_vector}) DESC APPROXIMATE LIMIT {top_k} - OPTIONS {{ METRIC: L2, TYPE: IVF, NPROBE: 8 }} - RETURN n.id AS id, euclidean(n.embedding, {gql_vector}) AS score + OPTIONS {{ METRIC: IP, TYPE: IVF, NPROBE: 8 }} + RETURN n.id AS id, inner_product(n.embedding, {gql_vector}) AS score """ try: - result = self.client.execute(gql) + result = self.execute_query(gql) except Exception as e: logger.error(f"[search_by_embedding] Query failed: {e}") return [] @@ -628,6 +794,7 @@ def search_by_embedding( values = row.values() id_val = values[0].as_string() score_val = values[1].as_double() + score_val = (score_val + 1) / 2 # align to neo4j, Normalized Cosine Score if threshold is None or score_val <= threshold: output.append({"id": id_val, "score": score_val}) return output @@ -635,7 +802,6 @@ def search_by_embedding( logger.error(f"[search_by_embedding] Result parse failed: {e}") return [] - # TODO def get_by_metadata(self, filters: list[dict[str, Any]]) -> list[str]: """ 1. ADD logic: "AND" vs "OR"(support logic combination); @@ -661,37 +827,48 @@ def get_by_metadata(self, filters: list[dict[str, Any]]) -> list[str]: - Can be used for faceted recall or prefiltering before embedding rerank. """ where_clauses = [] + + def _escape_value(value): + if isinstance(value, str): + return f'"{value}"' + elif isinstance(value, list): + return "[" + ", ".join(_escape_value(v) for v in value) + "]" + else: + return str(value) + for _i, f in enumerate(filters): field = f["field"] op = f.get("op", "=") value = f["value"] + escaped_value = _escape_value(value) + # Build WHERE clause if op == "=": - where_clauses.append(f"n.{field} = {value}") + where_clauses.append(f"n.{field} = {escaped_value}") elif op == "in": - where_clauses.append(f"n.{field} IN {value}") + where_clauses.append(f"n.{field} IN {escaped_value}") elif op == "contains": - where_clauses.append(f"ANY(x IN {value} WHERE x IN n.{field})") + where_clauses.append(f"ANY(x IN n.{field} WHERE x = {escaped_value})") elif op == "starts_with": - where_clauses.append(f"n.{field} STARTS WITH {value}") + where_clauses.append(f"n.{field} STARTS WITH {escaped_value}") elif op == "ends_with": - where_clauses.append(f"n.{field} ENDS WITH {value}") + where_clauses.append(f"n.{field} ENDS WITH {escaped_value}") elif op in [">", ">=", "<", "<="]: - where_clauses.append(f"n.{field} {op} {value}") + where_clauses.append(f"n.{field} {op} {escaped_value}") else: raise ValueError(f"Unsupported operator: {op}") - if not self.config.use_multi_db and self.config.user_name: - where_clauses.append(f"n.user_name = '{self.config.user_name}'") + if not self.config.use_multi_db and self.user_name: + where_clauses.append(f'n.user_name = "{self.config.user_name}"') where_str = " AND ".join(where_clauses) query = f"MATCH (n@Memory) WHERE {where_str} RETURN n.id AS id" try: - print("\n==========> query:\n", query) - result = self.client.execute(query) - return [record["id"].value for record in result] + result = self.execute_query(query) + ids = [record["id"].value for record in result] + return ids except Exception as e: logger.error(f"Failed to get metadata: {e}") @@ -740,7 +917,7 @@ def get_grouped_counts( group_by_fields = [] for field in group_fields: - alias = field.replace(".", "_") # 防止特殊字符 + alias = field.replace(".", "_") return_fields.append(f"n.{field} AS {alias}") group_by_fields.append(alias) # Full GQL query construction @@ -750,7 +927,7 @@ def get_grouped_counts( RETURN {", ".join(return_fields)}, COUNT(n) AS count GROUP BY {", ".join(group_by_fields)} """ - result = self.client.execute(gql) # Pure GQL string execution + result = self.execute_query(gql) # Pure GQL string execution output = [] for record in result: @@ -773,7 +950,7 @@ def clear(self) -> None: else: query = "MATCH (n) DETACH DELETE n" - self.client.execute(query) + self.execute_query(query) logger.info("Cleared all nodes from database.") except Exception as e: @@ -799,23 +976,20 @@ def export_graph(self) -> dict[str, Any]: try: full_node_query = f"{node_query} RETURN n" - node_result = self.client.execute(full_node_query) + node_result = self.execute_query(full_node_query) nodes = [] for row in node_result: node_wrapper = row.values()[0].as_node() props = node_wrapper.get_properties() - metadata = {key: self._parse_node(val) for key, val in props.items()} - - memory = metadata.get("memory", "") - - nodes.append({"id": node_wrapper.get_id(), "memory": memory, "metadata": metadata}) + node = self._parse_node(props) + nodes.append(node) except Exception as e: raise RuntimeError(f"[EXPORT GRAPH - NODES] Exception: {e}") from e try: full_edge_query = f"{edge_query} RETURN a.id AS source, b.id AS target, type(r) as edge" - edge_result = self.client.execute(full_edge_query) + edge_result = self.execute_query(full_edge_query) edges = [ { "source": row.values()[0].value, @@ -843,9 +1017,10 @@ def import_graph(self, data: dict[str, Any]) -> None: metadata["user_name"] = self.config.user_name metadata = _prepare_node_metadata(metadata) + metadata.update({"id": id, "memory": memory}) properties = ", ".join(f"{k}: {_format_value(v, k)}" for k, v in metadata.items()) node_gql = f"INSERT OR IGNORE (n@Memory {{{properties}}})" - self.client.execute(node_gql) + self.execute_query(node_gql) for edge in data.get("edges", []): source_id, target_id = edge["source"], edge["target"] @@ -857,7 +1032,7 @@ def import_graph(self, data: dict[str, Any]) -> None: MATCH (a@Memory {{id: "{source_id}"}}), (b@Memory {{id: "{target_id}"}}) INSERT OR IGNORE (a) -[e@{edge_type} {props}]-> (b) ''' - self.client.execute(edge_gql) + self.execute_query(edge_gql) def get_all_memory_items(self, scope: str) -> list[dict]: """ @@ -869,7 +1044,7 @@ def get_all_memory_items(self, scope: str) -> list[dict]: Returns: list[dict]: Full list of memory items under this scope. """ - if scope not in {"WorkingMemory", "LongTermMemory", "UserMemory"}: + if scope not in {"WorkingMemory", "LongTermMemory", "UserMemory", "OuterMemory"}: raise ValueError(f"Unsupported memory type scope: {scope}") where_clause = f"WHERE n.memory_type = '{scope}'" @@ -882,40 +1057,49 @@ def get_all_memory_items(self, scope: str) -> list[dict]: {where_clause} RETURN n """ + nodes = [] try: - results = self.client.execute(query) - return [self._parse_node(record["n"]) for record in results] + results = self.execute_query(query) + for rec in results: + node_props = rec["n"].as_node().get_properties() + nodes.append(self._parse_node(node_props)) except Exception as e: logger.error(f"Failed to get memories: {e}") + return nodes - # TODO def get_structure_optimization_candidates(self, scope: str) -> list[dict]: """ Find nodes that are likely candidates for structure optimization: - Isolated nodes, nodes with empty background, or nodes with exactly one child. - Plus: the child of any parent node that has exactly one child. """ - where_clause = f""" - WHERE n.memory_type = '{scope}' - AND n.status = 'activated' - AND NOT ( (n)-[r@PARENT]->() OR ()-[r@PARENT]->(n) ) - """ + where_clause = f''' + n.memory_type = "{scope}" + AND n.status = "activated" + ''' if not self.config.use_multi_db and self.config.user_name: - where_clause += f" AND n.user_name = '{self.config.user_name}'" + where_clause += f' AND n.user_name = "{self.config.user_name}"' query = f""" - MATCH (n@Memory) - {where_clause} - RETURN n.id AS id, n AS node - """ + USE `{self.db_name}` + MATCH (n@Memory) + WHERE {where_clause} + OPTIONAL MATCH (n)-[@PARENT]->(c@Memory) + OPTIONAL MATCH (p@Memory)-[@PARENT]->(n) + WHERE c IS NULL AND p IS NULL + RETURN n + """ + + candidates = [] try: - results = self.client.execute(query) - return [ - self._parse_node({"id": record["id"], **dict(record["node"])}) for record in results - ] + results = self.execute_query(query) + for rec in results: + node_props = rec["n"].as_node().get_properties() + candidates.append(self._parse_node(node_props)) except Exception as e: logger.error(f"Failed : {e}") + return candidates def drop_database(self) -> None: """ @@ -923,11 +1107,11 @@ def drop_database(self) -> None: WARNING: This operation is destructive and cannot be undone. """ if self.config.use_multi_db: - self.client.execute(f"DROP GRAPH {self.db_name}") - logger.info(f"Database '{self.db_name}' has been dropped.") + self.execute_query(f"DROP GRAPH `{self.db_name}`") + logger.info(f"Database '`{self.db_name}`' has been dropped.") else: raise ValueError( - f"Refusing to drop protected database: {self.db_name} in " + f"Refusing to drop protected database: `{self.db_name}` in " f"Shared Database Multi-Tenant mode" ) @@ -997,93 +1181,144 @@ def merge_nodes(self, id1: str, id2: str) -> str: def _ensure_database_exists(self): create_tag = """ - CREATE GRAPH TYPE IF NOT EXISTS MemoryGraphType AS { + CREATE GRAPH TYPE IF NOT EXISTS MemOSType AS { NODE Memory (:MemoryTag { id STRING, memory STRING, user_name STRING, - created_at STRING, - updated_at STRING, + user_id STRING, + session_id STRING, status STRING, - node_type STRING, - memory_time STRING, - source STRING, + key STRING, confidence FLOAT, - entities LIST, tags LIST, - visibility STRING, + created_at STRING, + updated_at STRING, memory_type STRING, - key STRING, sources LIST, + source STRING, + node_type STRING, + visibility STRING, usage LIST, background STRING, - hierarchy_level STRING, embedding VECTOR<3072, FLOAT>, PRIMARY KEY(id) }), EDGE RELATE_TO (Memory) -[{user_name STRING}]-> (Memory), - EDGE PARENT (Memory) -[{user_name STRING}]-> (Memory) + EDGE PARENT (Memory) -[{user_name STRING}]-> (Memory), + EDGE AGGREGATE_TO (Memory) -[{user_name STRING}]-> (Memory), + EDGE MERGED_TO (Memory) -[{user_name STRING}]-> (Memory), + EDGE INFERS (Memory) -[{user_name STRING}]-> (Memory), + EDGE FOLLOWS (Memory) -[{user_name STRING}]-> (Memory) } """ - create_graph = "CREATE GRAPH IF NOT EXISTS memory_graph TYPED MemoryGraphType" - set_graph_working = "SESSION SET GRAPH memory_graph" + create_graph = f"CREATE GRAPH IF NOT EXISTS `{self.db_name}` TYPED MemOSType" + set_graph_working = f"SESSION SET GRAPH `{self.db_name}`" try: - self.client.execute(create_tag) - self.client.execute(create_graph) - self.client.execute(set_graph_working) - logger.info("✅ Graph `memory_graph` is now the working graph.") + self.execute_query(create_tag, auto_set_db=False) + self.execute_query(create_graph, auto_set_db=False) + self.execute_query(set_graph_working) + logger.info(f"✅ Graph ``{self.db_name}`` is now the working graph.") except Exception as e: - logger.error(f"❌ Failed to create tag: {e}") + logger.error(f"❌ Failed to create tag: {e} trace: {traceback.format_exc()}") - # TODO - def _vector_index_exists(self, index_name: str = "memory_vector_index") -> bool: - raise NotImplementedError - - # TODO def _create_vector_index( self, label: str, vector_property: str, dimensions: int, index_name: str ) -> None: """ Create a vector index for the specified property in the label. """ - raise NotImplementedError + create_vector_index = f""" + CREATE VECTOR INDEX IF NOT EXISTS {index_name} + ON NODE Memory::{vector_property} + OPTIONS {{ + DIM: {dimensions}, + METRIC: IP, + TYPE: IVF, + NLIST: 100, + TRAINSIZE: 1000 + }} + FOR `{self.db_name}` + """ + self.execute_query(create_vector_index) - # TODO def _create_basic_property_indexes(self) -> None: """ - Create standard B-tree indexes on memory_type, created_at, + Create standard B-tree indexes on status, memory_type, created_at and updated_at fields. Create standard B-tree indexes on user_name when use Shared Database - Multi-Tenant Mode + Multi-Tenant Mode. """ - raise NotImplementedError + fields = ["status", "memory_type", "created_at", "updated_at"] + if not self.config.use_multi_db: + fields.append("user_name") + + for field in fields: + index_name = f"idx_memory_{field}" + gql = f""" + CREATE INDEX IF NOT EXISTS {index_name} ON NODE Memory({field}) + FOR `{self.db_name}` + """ + try: + self.execute_query(gql) + logger.info(f"✅ Created index: {index_name} on field {field}") + except Exception as e: + logger.error(f"❌ Failed to create index {index_name}: {e}") def _index_exists(self, index_name: str) -> bool: """ Check if an index with the given name exists. """ - raise NotImplementedError + """ + Check if a vector index with the given name exists in NebulaGraph. + + Args: + index_name (str): The name of the index to check. + + Returns: + bool: True if the index exists, False otherwise. + """ + query = "SHOW VECTOR INDEXES" + try: + result = self.execute_query(query) + return any(row.values()[0].as_string() == index_name for row in result) + except Exception as e: + logger.error(f"[Nebula] Failed to check index existence: {e}") + return False - def _parse_node(self, value: ValueWrapper) -> Any: - if value is None or value.is_null(): + def _parse_value(self, value: Any) -> Any: + """turn Nebula ValueWrapper to Python type""" + from nebulagraph_python.value_wrapper import ValueWrapper + + if value is None or (hasattr(value, "is_null") and value.is_null()): return None try: - primitive_value = value.cast_primitive() + prim = value.cast_primitive() if isinstance(value, ValueWrapper) else value except Exception as e: - logger.warning(f"cast_primitive failed for value: {value}, error: {e}") - try: - primitive_value = value.cast() - except Exception as e2: - logger.warning(f"cast failed for value: {value}, error: {e2}") - return str(value) + logger.warning(f"Error when decode Nebula ValueWrapper: {e}") + prim = value.cast() if isinstance(value, ValueWrapper) else value - if isinstance(primitive_value, ValueWrapper): - return self._parse_node(primitive_value) + if isinstance(prim, ValueWrapper): + return self._parse_value(prim) + if isinstance(prim, list): + return [self._parse_value(v) for v in prim] + if type(prim).__name__ == "NVector": + return list(prim.values) - if isinstance(primitive_value, list): - return [ - self._parse_node(v) if isinstance(v, ValueWrapper) else v for v in primitive_value - ] + return prim # already a Python primitive + + def _parse_node(self, props: dict[str, Any]) -> dict[str, Any]: + parsed = {k: self._parse_value(v) for k, v in props.items()} + + for tf in ("created_at", "updated_at"): + if tf in parsed and hasattr(parsed[tf], "isoformat"): + parsed[tf] = parsed[tf].isoformat() + + node_id = parsed.pop("id") + memory = parsed.pop("memory", "") + parsed.pop("user_name", None) + metadata = parsed + metadata["type"] = metadata.pop("node_type") - return primitive_value + return {"id": node_id, "memory": memory, "metadata": metadata} diff --git a/src/memos/graph_dbs/neo4j.py b/src/memos/graph_dbs/neo4j.py index 1489ff541..c22c2666e 100644 --- a/src/memos/graph_dbs/neo4j.py +++ b/src/memos/graph_dbs/neo4j.py @@ -114,14 +114,14 @@ def get_memory_count(self, memory_type: str) -> int: ) return result.single()["count"] - def count_nodes(self, scope: str) -> int: + def node_not_exist(self, scope: str) -> int: query = """ MATCH (n:Memory) WHERE n.memory_type = $scope """ if not self.config.use_multi_db and self.config.user_name: query += "\nAND n.user_name = $user_name" - query += "\nRETURN count(n) AS count" + query += "\nRETURN n LIMIT 1" with self.driver.session(database=self.db_name) as session: result = session.run( @@ -131,7 +131,7 @@ def count_nodes(self, scope: str) -> int: "user_name": self.config.user_name if self.config.user_name else None, }, ) - return result.single()["count"] + return result.single() is None def remove_oldest_memory(self, memory_type: str, keep_latest: int) -> None: """ @@ -920,7 +920,7 @@ def get_all_memory_items(self, scope: str) -> list[dict]: Returns: list[dict]: Full list of memory items under this scope. """ - if scope not in {"WorkingMemory", "LongTermMemory", "UserMemory"}: + if scope not in {"WorkingMemory", "LongTermMemory", "UserMemory", "OuterMemory"}: raise ValueError(f"Unsupported memory type scope: {scope}") where_clause = "WHERE n.memory_type = $scope" diff --git a/src/memos/log.py b/src/memos/log.py index 12a92b5b5..3ba0fe485 100644 --- a/src/memos/log.py +++ b/src/memos/log.py @@ -48,7 +48,7 @@ def _setup_logfile() -> Path: "class": "logging.handlers.RotatingFileHandler", "filename": _setup_logfile(), "maxBytes": 1024**2 * 10, - "backupCount": 3, + "backupCount": 10, "formatter": "standard", }, }, diff --git a/src/memos/mem_cube/utils.py b/src/memos/mem_cube/utils.py index 0e7afaf39..7d5414b0f 100644 --- a/src/memos/mem_cube/utils.py +++ b/src/memos/mem_cube/utils.py @@ -71,10 +71,6 @@ def merge_config_with_default( # Define graph_db fields to preserve (user-specific) preserve_graph_fields = { - "uri", - "user", - "password", - "db_name", "auto_create", "user_name", "use_multi_db", diff --git a/src/memos/mem_os/core.py b/src/memos/mem_os/core.py index 0dcb72d48..7538db4ee 100644 --- a/src/memos/mem_os/core.py +++ b/src/memos/mem_os/core.py @@ -17,6 +17,7 @@ from memos.mem_scheduler.schemas.general_schemas import ( ADD_LABEL, ANSWER_LABEL, + QUERY_LABEL, ) from memos.mem_scheduler.schemas.message_schemas import ScheduleMessageItem from memos.mem_user.user_manager import UserManager, UserRole @@ -167,6 +168,14 @@ def mem_reorganizer_off(self) -> bool: if mem_cube.text_mem and mem_cube.text_mem.is_reorganize: logger.info(f"close reorganizer for {mem_cube.text_mem.config.cube_id}") mem_cube.text_mem.memory_manager.close() + mem_cube.text_mem.memory_manager.wait_reorganizer() + + def mem_reorganizer_wait(self) -> bool: + for mem_cube in self.mem_cubes.values(): + logger.info(f"try to close reorganizer for {mem_cube.text_mem.config.cube_id}") + if mem_cube.text_mem and mem_cube.text_mem.is_reorganize: + logger.info(f"close reorganizer for {mem_cube.text_mem.config.cube_id}") + mem_cube.text_mem.memory_manager.wait_reorganizer() def _register_chat_history(self, user_id: str | None = None) -> None: """Initialize chat history with user ID.""" @@ -267,7 +276,7 @@ def chat(self, query: str, user_id: str | None = None, base_prompt: str | None = user_id=target_user_id, mem_cube_id=mem_cube_id, mem_cube=mem_cube, - label=ADD_LABEL, + label=QUERY_LABEL, content=query, timestamp=datetime.now(), ) @@ -521,6 +530,8 @@ def search( user_id: str | None = None, install_cube_ids: list[str] | None = None, top_k: int | None = None, + mode: Literal["fast", "fine"] = "fast", + internet_search: bool = False, ) -> MOSSearchResult: """ Search for textual memories across all registered MemCubes. @@ -558,7 +569,11 @@ def search( and self.config.enable_textual_memory ): memories = mem_cube.text_mem.search( - query, top_k=top_k if top_k else self.config.top_k + query, + top_k=top_k if top_k else self.config.top_k, + mode=mode, + manual_close_internet=not internet_search, + info={"user_id": target_user_id, "session_id": str(uuid.uuid4())}, ) result["text_mem"].append({"cube_id": mem_cube_id, "memories": memories}) logger.info( @@ -631,6 +646,9 @@ def add( for mem in memories: mem_id_list: list[str] = self.mem_cubes[mem_cube_id].text_mem.add(mem) mem_ids.extend(mem_id_list) + logger.info( + f"Added memory user {target_user_id} to memcube {mem_cube_id}: {mem_id_list}" + ) # submit messages for scheduler if self.enable_mem_scheduler and self.mem_scheduler is not None: @@ -671,6 +689,9 @@ def add( mem_ids = [] for mem in memories: mem_id_list: list[str] = self.mem_cubes[mem_cube_id].text_mem.add(mem) + logger.info( + f"Added memory user {target_user_id} to memcube {mem_cube_id}: {mem_id_list}" + ) mem_ids.extend(mem_id_list) # submit messages for scheduler diff --git a/src/memos/mem_os/product.py b/src/memos/mem_os/product.py index 31ab55a99..049142691 100644 --- a/src/memos/mem_os/product.py +++ b/src/memos/mem_os/product.py @@ -23,7 +23,11 @@ remove_embedding_recursive, sort_children_by_memory_type, ) -from memos.mem_scheduler.schemas import ANSWER_LABEL, QUERY_LABEL, ScheduleMessageItem +from memos.mem_scheduler.schemas.general_schemas import ( + ANSWER_LABEL, + QUERY_LABEL, +) +from memos.mem_scheduler.schemas.message_schemas import ScheduleMessageItem from memos.mem_user.persistent_user_manager import PersistentUserManager from memos.mem_user.user_manager import UserRole from memos.memories.textual.item import ( @@ -48,8 +52,10 @@ class MOSProduct(MOSCore): def __init__( self, default_config: MOSConfig | None = None, - max_user_instances: int = 100, + max_user_instances: int = 1, default_cube_config: GeneralMemCubeConfig | None = None, + online_bot=None, + error_bot=None, ): """ Initialize MOSProduct with an optional default configuration. @@ -58,6 +64,8 @@ def __init__( default_config (MOSConfig | None): Default configuration for new users max_user_instances (int): Maximum number of user instances to keep in memory default_cube_config (GeneralMemCubeConfig | None): Default cube configuration for loading cubes + online_bot: DingDing online_bot function or None if disabled + error_bot: DingDing error_bot function or None if disabled """ # Initialize with a root config for shared resources if default_config is None: @@ -84,6 +92,8 @@ def __init__( self.default_config = default_config self.default_cube_config = default_cube_config self.max_user_instances = max_user_instances + self.online_bot = online_bot + self.error_bot = error_bot # User-specific data structures self.user_configs: dict[str, MOSConfig] = {} @@ -359,7 +369,7 @@ def _build_system_prompt(self, user_id: str, memories_all: list[TextualMemoryIte for i, memory in enumerate(memories_all, 1): # Format: [memory_id]: memory_content memory_id = f"{memory.id.split('-')[0]}" if hasattr(memory, "id") else f"mem_{i}" - memory_content = memory.memory if hasattr(memory, "memory") else str(memory) + memory_content = memory.memory[:500] if hasattr(memory, "memory") else str(memory) memory_context += f"{memory_id}: {memory_content}\n" return base_prompt + memory_context @@ -419,27 +429,39 @@ def _process_streaming_references_complete(self, text_buffer: str) -> tuple[str, # No reference tags found, return all text return text_buffer, "" - def _extract_references_from_response(self, response: str) -> list[dict]: + def _extract_references_from_response(self, response: str) -> tuple[str, list[dict]]: """ - Extract reference information from the response. + Extract reference information from the response and return clean text. Args: response (str): The complete response text. Returns: - list[dict]: List of reference information. + tuple[str, list[dict]]: A tuple containing: + - clean_text: Text with reference markers removed + - references: List of reference information """ import re - references = [] - # Pattern to match [refid:memoriesID] - pattern = r"\[(\d+):([^\]]+)\]" + try: + references = [] + # Pattern to match [refid:memoriesID] + pattern = r"\[(\d+):([^\]]+)\]" + + matches = re.findall(pattern, response) + for ref_number, memory_id in matches: + references.append({"memory_id": memory_id, "reference_number": int(ref_number)}) - matches = re.findall(pattern, response) - for ref_number, memory_id in matches: - references.append({"memory_id": memory_id, "reference_number": int(ref_number)}) + # Remove all reference markers from the text to get clean text + clean_text = re.sub(pattern, "", response) - return references + # Clean up any extra whitespace that might be left after removing markers + clean_text = re.sub(r"\s+", " ", clean_text).strip() + + return clean_text, references + except Exception as e: + logger.error(f"Error extracting references from response: {e}", exc_info=True) + return response, [] def _chunk_response_with_tiktoken( self, response: str, chunk_size: int = 5 @@ -494,6 +516,14 @@ def _send_message_to_scheduler( ) self.mem_scheduler.submit_messages(messages=[message_item]) + def _filter_memories_by_threshold( + self, memories: list[TextualMemoryItem], threshold: float = 0.20 + ) -> list[TextualMemoryItem]: + """ + Filter memories by threshold. + """ + return [memory for memory in memories if memory.metadata.relativity >= threshold] + def register_mem_cube( self, mem_cube_name_or_path_or_object: str | GeneralMemCube, @@ -601,7 +631,7 @@ def user_register( try: default_mem_cube.dump(mem_cube_name_or_path) except Exception as e: - print(e) + logger.error(f"Failed to dump default cube: {e}") # Register the default cube with MOS self.register_mem_cube( @@ -670,7 +700,7 @@ def get_suggestion_query(self, user_id: str, language: str = "zh") -> list[str]: "text_mem" ] if text_mem_result: - memories = "\n".join([m.memory for m in text_mem_result[0]["memories"]]) + memories = "\n".join([m.memory[:200] for m in text_mem_result[0]["memories"]]) else: memories = "" message_list = [{"role": "system", "content": suggestion_prompt.format(memories=memories)}] @@ -679,63 +709,14 @@ def get_suggestion_query(self, user_id: str, language: str = "zh") -> list[str]: response_json = json.loads(clean_response) return response_json["query"] - def chat( - self, - query: str, - user_id: str, - cube_id: str | None = None, - history: MessageList | None = None, - ) -> Generator[str, None, None]: - """Chat with LLM SSE Type. - Args: - query (str): Query string. - user_id (str): User ID. - cube_id (str, optional): Custom cube ID for user. - history (list[dict], optional): Chat history. - - Returns: - Generator[str, None, None]: The response string generator. - """ - # Use MOSCore's built-in validation - if cube_id: - self._validate_cube_access(user_id, cube_id) - else: - self._validate_user_exists(user_id) - - # Load user cubes if not already loaded - self._load_user_cubes(user_id, self.default_cube_config) - time_start = time.time() - memories_list = super().search(query, user_id)["text_mem"] - # Get response from parent MOSCore (returns string, not generator) - response = super().chat(query, user_id) - time_end = time.time() - - # Use tiktoken for proper token-based chunking - for chunk in self._chunk_response_with_tiktoken(response, chunk_size=5): - chunk_data = f"data: {json.dumps({'type': 'text', 'content': chunk})}\n\n" - yield chunk_data - - # Prepare reference data - reference = [] - for memories in memories_list: - memories_json = memories.model_dump() - memories_json["metadata"]["ref_id"] = f"[{memories.id.split('-')[0]}]" - memories_json["metadata"]["embedding"] = [] - memories_json["metadata"]["sources"] = [] - reference.append(memories_json) - - yield f"data: {json.dumps({'type': 'reference', 'content': reference})}\n\n" - total_time = round(float(time_end - time_start), 1) - - yield f"data: {json.dumps({'type': 'time', 'content': {'total_time': total_time, 'speed_improvement': '23%'}})}\n\n" - yield f"data: {json.dumps({'type': 'end'})}\n\n" - def chat_with_references( self, query: str, user_id: str, cube_id: str | None = None, history: MessageList | None = None, + top_k: int = 10, + internet_search: bool = False, ) -> Generator[str, None, None]: """ Chat with LLM with memory references and streaming output. @@ -751,15 +732,24 @@ def chat_with_references( """ self._load_user_cubes(user_id, self.default_cube_config) - time_start = time.time() memories_list = [] + yield f"data: {json.dumps({'type': 'status', 'data': '0'})}\n\n" memories_result = super().search( - query, user_id, install_cube_ids=[cube_id] if cube_id else None, top_k=10 + query, + user_id, + install_cube_ids=[cube_id] if cube_id else None, + top_k=top_k, + mode="fine", + internet_search=internet_search, )["text_mem"] + yield f"data: {json.dumps({'type': 'status', 'data': '1'})}\n\n" + self._send_message_to_scheduler( + user_id=user_id, mem_cube_id=cube_id, query=query, label=QUERY_LABEL + ) if memories_result: memories_list = memories_result[0]["memories"] - + memories_list = self._filter_memories_by_threshold(memories_list) # Build custom system prompt with relevant memories system_prompt = self._build_system_prompt(user_id, memories_list) @@ -768,12 +758,14 @@ def chat_with_references( self._register_chat_history(user_id) chat_history = self.chat_history_manager[user_id] + if history: + chat_history.chat_history = history[-10:] current_messages = [ {"role": "system", "content": system_prompt}, *chat_history.chat_history, {"role": "user", "content": query}, ] - + yield f"data: {json.dumps({'type': 'status', 'data': '2'})}\n\n" # Generate response with custom prompt past_key_values = None response_stream = None @@ -809,7 +801,7 @@ def chat_with_references( # Initialize buffer for streaming buffer = "" full_response = "" - + token_count = 0 # Use tiktoken for proper token-based chunking if self.config.chat_model.backend not in ["huggingface", "vllm"]: # For non-huggingface backends, we need to collect the full response first @@ -822,6 +814,7 @@ def chat_with_references( for chunk in response_stream: if chunk in ["", ""]: continue + token_count += 1 buffer += chunk full_response += chunk @@ -848,22 +841,60 @@ def chat_with_references( memories_json["metadata"]["embedding"] = [] memories_json["metadata"]["sources"] = [] memories_json["metadata"]["memory"] = memories.memory + memories_json["metadata"]["id"] = memories.id reference.append({"metadata": memories_json["metadata"]}) yield f"data: {json.dumps({'type': 'reference', 'data': reference})}\n\n" + # set kvcache improve speed + speed_improvement = round(float((len(system_prompt) / 2) * 0.0048 + 44.5), 1) total_time = round(float(time_end - time_start), 1) - yield f"data: {json.dumps({'type': 'time', 'data': {'total_time': total_time, 'speed_improvement': '23%'}})}\n\n" - chat_history.chat_history.append({"role": "user", "content": query}) - chat_history.chat_history.append({"role": "assistant", "content": full_response}) - self._send_message_to_scheduler( - user_id=user_id, mem_cube_id=cube_id, query=query, label=QUERY_LABEL - ) - self._send_message_to_scheduler( - user_id=user_id, mem_cube_id=cube_id, query=full_response, label=ANSWER_LABEL - ) - self.chat_history_manager[user_id] = chat_history + yield f"data: {json.dumps({'type': 'time', 'data': {'total_time': total_time, 'speed_improvement': f'{speed_improvement}%'}})}\n\n" yield f"data: {json.dumps({'type': 'end'})}\n\n" + + logger.info(f"user_id: {user_id}, cube_id: {cube_id}, current_messages: {current_messages}") + logger.info(f"user_id: {user_id}, cube_id: {cube_id}, full_response: {full_response}") + + clean_response, extracted_references = self._extract_references_from_response(full_response) + logger.info(f"Extracted {len(extracted_references)} references from response") + + # Send chat report if online_bot is available + try: + from memos.memos_tools.notification_utils import send_online_bot_notification + + # Prepare data for online_bot + chat_data = { + "query": query, + "user_id": user_id, + "cube_id": cube_id, + "system_prompt": system_prompt, + "full_response": full_response, + } + + system_data = { + "references": extracted_references, + "time_start": time_start, + "time_end": time_end, + "speed_improvement": speed_improvement, + } + + emoji_config = {"chat": "💬", "system_info": "📊"} + + send_online_bot_notification( + online_bot=self.online_bot, + header_name="MemOS Chat Report", + sub_title_name="chat_with_references", + title_color="#00956D", + other_data1=chat_data, + other_data2=system_data, + emoji=emoji_config, + ) + except Exception as e: + logger.warning(f"Failed to send chat notification: {e}") + + self._send_message_to_scheduler( + user_id=user_id, mem_cube_id=cube_id, query=clean_response, label=ANSWER_LABEL + ) self.add( user_id=user_id, messages=[ @@ -874,18 +905,12 @@ def chat_with_references( }, { "role": "assistant", - "content": full_response, + "content": clean_response, # Store clean text without reference markers "chat_time": str(datetime.now().strftime("%Y-%m-%d %H:%M:%S")), }, ], mem_cube_id=cube_id, ) - # Keep chat history under 30 messages by removing oldest conversation pair - if len(self.chat_history_manager[user_id].chat_history) > 10: - self.chat_history_manager[user_id].chat_history.pop(0) # Remove oldest user message - self.chat_history_manager[user_id].chat_history.pop( - 0 - ) # Remove oldest assistant response def get_all( self, @@ -988,6 +1013,7 @@ def get_subgraph( user_id: str, query: str, mem_cube_ids: list[str] | None = None, + top_k: int = 20, ) -> list[dict[str, Any]]: """Get all memory items for a user. @@ -1003,7 +1029,7 @@ def get_subgraph( # Load user cubes if not already loaded self._load_user_cubes(user_id, self.default_cube_config) memory_list = self._get_subgraph( - query=query, mem_cube_id=mem_cube_ids[0], user_id=user_id, top_k=20 + query=query, mem_cube_id=mem_cube_ids[0], user_id=user_id, top_k=top_k )["text_mem"] reformat_memory_list = [] for memory in memory_list: @@ -1030,15 +1056,18 @@ def get_subgraph( return reformat_memory_list def search( - self, query: str, user_id: str, install_cube_ids: list[str] | None = None, top_k: int = 20 + self, + query: str, + user_id: str, + install_cube_ids: list[str] | None = None, + top_k: int = 10, + mode: Literal["fast", "fine"] = "fast", ): """Search memories for a specific user.""" - # Validate user access - self._validate_user_access(user_id) # Load user cubes if not already loaded self._load_user_cubes(user_id, self.default_cube_config) - search_result = super().search(query, user_id, install_cube_ids, top_k) + search_result = super().search(query, user_id, install_cube_ids, top_k, mode=mode) text_memory_list = search_result["text_mem"] reformat_memory_list = [] for memory in text_memory_list: diff --git a/src/memos/mem_os/utils/format_utils.py b/src/memos/mem_os/utils/format_utils.py index 8465b44f0..049cd6995 100644 --- a/src/memos/mem_os/utils/format_utils.py +++ b/src/memos/mem_os/utils/format_utils.py @@ -533,7 +533,7 @@ def convert_graph_to_tree_forworkmem( node_name = extract_node_name(memory) memory_key = node.get("metadata", {}).get("key", node_name) usage = node.get("metadata", {}).get("usage", []) - frequency = len(usage) + frequency = len(usage) if len(usage) < 100 else 100 node_map[node["id"]] = { "id": node["id"], "value": memory, diff --git a/src/memos/mem_reader/simple_struct.py b/src/memos/mem_reader/simple_struct.py index 8a01e9316..2070b5046 100644 --- a/src/memos/mem_reader/simple_struct.py +++ b/src/memos/mem_reader/simple_struct.py @@ -180,8 +180,12 @@ def get_scene_data_info(self, scene_data: list, type: str) -> list[str]: elif type == "doc": for item in scene_data: try: - parsed_text = parser.parse(item) - results.append({"file": item, "text": parsed_text}) + if not isinstance(item, str): + parsed_text = parser.parse(item) + results.append({"file": "pure_text", "text": parsed_text}) + else: + parsed_text = item + results.append({"file": item, "text": parsed_text}) except Exception as e: print(f"Error parsing file {item}: {e!s}") diff --git a/src/memos/mem_scheduler/base_scheduler.py b/src/memos/mem_scheduler/base_scheduler.py index 0335bd5e4..295090f93 100644 --- a/src/memos/mem_scheduler/base_scheduler.py +++ b/src/memos/mem_scheduler/base_scheduler.py @@ -20,7 +20,9 @@ DEFAULT_ACT_MEM_DUMP_PATH, DEFAULT_CONSUME_INTERVAL_SECONDS, DEFAULT_THREAD__POOL_MAX_WORKERS, + MemCubeID, TreeTextMemory_SEARCH_METHOD, + UserID, ) from memos.mem_scheduler.schemas.message_schemas import ( ScheduleLogForWebItem, @@ -81,7 +83,7 @@ def __init__(self, config: BaseSchedulerConfig): # other attributes self._context_lock = threading.Lock() - self._current_user_id: str | None = None + self.current_user_id: UserID | str | None = None self.auth_config_path: str | Path | None = self.config.get("auth_config_path", None) self.auth_config = None self.rabbitmq_config = None @@ -113,20 +115,20 @@ def initialize_modules(self, chat_llm: BaseLLM, process_llm: BaseLLM | None = No @property def mem_cube(self) -> GeneralMemCube: """The memory cube associated with this MemChat.""" - return self._current_mem_cube + return self.current_mem_cube @mem_cube.setter def mem_cube(self, value: GeneralMemCube) -> None: """The memory cube associated with this MemChat.""" - self._current_mem_cube = value + self.current_mem_cube = value self.retriever.mem_cube = value def _set_current_context_from_message(self, msg: ScheduleMessageItem) -> None: """Update current user/cube context from the incoming message (thread-safe).""" with self._context_lock: - self._current_user_id = msg.user_id - self._current_mem_cube_id = msg.mem_cube_id - self._current_mem_cube = msg.mem_cube + self.current_user_id = msg.user_id + self.current_mem_cube_id = msg.mem_cube_id + self.current_mem_cube = msg.mem_cube def transform_memories_to_monitors( self, memories: list[TextualMemoryItem] @@ -181,9 +183,8 @@ def transform_memories_to_monitors( def replace_working_memory( self, - queries: list[str], - user_id: str, - mem_cube_id: str, + user_id: UserID | str, + mem_cube_id: MemCubeID | str, mem_cube: GeneralMemCube, original_memory: list[TextualMemoryItem], new_memory: list[TextualMemoryItem], @@ -194,10 +195,10 @@ def replace_working_memory( text_mem_base: TreeTextMemory = text_mem_base # process rerank memories with llm - quey_history = self.monitor.query_monitors.get_queries_with_timesort() + query_history = self.monitor.query_monitors.get_queries_with_timesort() memories_with_new_order, rerank_success_flag = ( self.retriever.process_and_rerank_memories( - queries=quey_history, + queries=query_history, original_memory=original_memory, new_memory=new_memory, top_k=self.top_k, @@ -246,8 +247,8 @@ def replace_working_memory( def initialize_working_memory_monitors( self, - user_id: str, - mem_cube_id: str, + user_id: UserID | str, + mem_cube_id: MemCubeID | str, mem_cube: GeneralMemCube, ): text_mem_base: TreeTextMemory = mem_cube.text_mem @@ -267,8 +268,8 @@ def update_activation_memory( self, new_memories: list[str | TextualMemoryItem], label: str, - user_id: str, - mem_cube_id: str, + user_id: UserID | str, + mem_cube_id: MemCubeID | str, mem_cube: GeneralMemCube, ) -> None: """ @@ -344,60 +345,70 @@ def update_activation_memory_periodically( self, interval_seconds: int, label: str, - user_id: str, - mem_cube_id: str, + user_id: UserID | str, + mem_cube_id: MemCubeID | str, mem_cube: GeneralMemCube, ): - new_activation_memories = [] + try: + if ( + self.monitor.last_activation_mem_update_time == datetime.min + or self.monitor.timed_trigger( + last_time=self.monitor.last_activation_mem_update_time, + interval_seconds=interval_seconds, + ) + ): + logger.info( + f"Updating activation memory for user {user_id} and mem_cube {mem_cube_id}" + ) - if self.monitor.timed_trigger( - last_time=self.monitor.last_activation_mem_update_time, - interval_seconds=interval_seconds, - ): - logger.info(f"Updating activation memory for user {user_id} and mem_cube {mem_cube_id}") + if ( + user_id not in self.monitor.working_memory_monitors + or mem_cube_id not in self.monitor.working_memory_monitors[user_id] + or len(self.monitor.working_memory_monitors[user_id][mem_cube_id].memories) == 0 + ): + logger.warning( + "No memories found in working_memory_monitors, initializing from current working_memories" + ) + self.initialize_working_memory_monitors( + user_id=user_id, + mem_cube_id=mem_cube_id, + mem_cube=mem_cube, + ) - if len(self.monitor.working_memory_monitors[user_id][mem_cube_id].memories) == 0: - logger.warning( - "No memories found in working_memory_monitors, initializing from current working_memories" + self.monitor.update_activation_memory_monitors( + user_id=user_id, mem_cube_id=mem_cube_id, mem_cube=mem_cube ) - self.initialize_working_memory_monitors( + + new_activation_memories = [ + m.memory_text + for m in self.monitor.activation_memory_monitors[user_id][mem_cube_id].memories + ] + + logger.info( + f"Collected {len(new_activation_memories)} new memory entries for processing" + ) + + self.update_activation_memory( + new_memories=new_activation_memories, + label=label, user_id=user_id, mem_cube_id=mem_cube_id, mem_cube=mem_cube, ) - self.monitor.update_activation_memory_monitors( - user_id=user_id, mem_cube_id=mem_cube_id, mem_cube=mem_cube - ) - - new_activation_memories = [ - m.memory_text - for m in self.monitor.activation_memory_monitors[user_id][mem_cube_id].memories - ] - - logger.info( - f"Collected {len(new_activation_memories)} new memory entries for processing" - ) - - self.update_activation_memory( - new_memories=new_activation_memories, - label=label, - user_id=user_id, - mem_cube_id=mem_cube_id, - mem_cube=mem_cube, - ) - - self.monitor.last_activation_mem_update_time = datetime.now() + self.monitor.last_activation_mem_update_time = datetime.now() - logger.debug( - f"Activation memory update completed at {self.monitor.last_activation_mem_update_time}" - ) - else: - logger.info( - f"Skipping update - {interval_seconds} second interval not yet reached. " - f"Last update time is {self.monitor.last_activation_mem_update_time} and now is" - f"{datetime.now()}" - ) + logger.debug( + f"Activation memory update completed at {self.monitor.last_activation_mem_update_time}" + ) + else: + logger.info( + f"Skipping update - {interval_seconds} second interval not yet reached. " + f"Last update time is {self.monitor.last_activation_mem_update_time} and now is" + f"{datetime.now()}" + ) + except Exception as e: + logger.error(f"Error: {e}", exc_info=True) def submit_messages(self, messages: ScheduleMessageItem | list[ScheduleMessageItem]): """Submit multiple messages to the message queue.""" diff --git a/src/memos/mem_scheduler/general_scheduler.py b/src/memos/mem_scheduler/general_scheduler.py index f4de9e8fb..d0b2d61e4 100644 --- a/src/memos/mem_scheduler/general_scheduler.py +++ b/src/memos/mem_scheduler/general_scheduler.py @@ -9,6 +9,8 @@ ANSWER_LABEL, DEFAULT_MAX_QUERY_KEY_WORDS, QUERY_LABEL, + MemCubeID, + UserID, ) from memos.mem_scheduler.schemas.message_schemas import ScheduleMessageItem from memos.mem_scheduler.schemas.monitor_schemas import QueryMonitorItem @@ -31,6 +33,51 @@ def __init__(self, config: GeneralSchedulerConfig): } self.dispatcher.register_handlers(handlers) + # for evaluation + def search_for_eval( + self, + query: str, + user_id: UserID | str, + top_k: int, + ) -> list[str]: + query_keywords = self.monitor.extract_query_keywords(query=query) + logger.info(f'Extract keywords "{query_keywords}" from query "{query}"') + + item = QueryMonitorItem( + query_text=query, + keywords=query_keywords, + max_keywords=DEFAULT_MAX_QUERY_KEY_WORDS, + ) + self.monitor.query_monitors.put(item=item) + logger.debug( + f"Queries in monitor are {self.monitor.query_monitors.get_queries_with_timesort()}." + ) + + queries = [query] + + # recall + cur_working_memory, new_candidates = self.process_session_turn( + queries=queries, + user_id=user_id, + mem_cube_id=self.current_mem_cube_id, + mem_cube=self.current_mem_cube, + top_k=self.top_k, + ) + logger.info(f"Processed {queries} and get {len(new_candidates)} new candidate memories.") + + # rerank + new_order_working_memory = self.replace_working_memory( + user_id=user_id, + mem_cube_id=self.current_mem_cube_id, + mem_cube=self.current_mem_cube, + original_memory=cur_working_memory, + new_memory=new_candidates, + ) + new_order_working_memory = new_order_working_memory[:top_k] + logger.info(f"size of new_order_working_memory: {len(new_order_working_memory)}") + + return [m.memory for m in new_order_working_memory] + def _query_message_consumer(self, messages: list[ScheduleMessageItem]) -> None: """ Process and handle query trigger messages from the queue. @@ -88,7 +135,6 @@ def _query_message_consumer(self, messages: list[ScheduleMessageItem]) -> None: # rerank new_order_working_memory = self.replace_working_memory( - queries=queries, user_id=user_id, mem_cube_id=mem_cube_id, mem_cube=mem_cube, @@ -119,7 +165,7 @@ def _answer_message_consumer(self, messages: list[ScheduleMessageItem]) -> None: # for status update self._set_current_context_from_message(msg=messages[0]) - # update acivation memories + # update activation memories if self.enable_act_memory_update: if ( len(self.monitor.working_memory_monitors[user_id][mem_cube_id].memories) @@ -145,49 +191,56 @@ def _add_message_consumer(self, messages: list[ScheduleMessageItem]) -> None: grouped_messages = self.dispatcher.group_messages_by_user_and_cube(messages=messages) self.validate_schedule_messages(messages=messages, label=ADD_LABEL) - - for user_id in grouped_messages: - for mem_cube_id in grouped_messages[user_id]: - messages = grouped_messages[user_id][mem_cube_id] - if len(messages) == 0: - return - - # for status update - self._set_current_context_from_message(msg=messages[0]) - - # submit logs - for msg in messages: - userinput_memory_ids = json.loads(msg.content) - mem_cube = msg.mem_cube - for memory_id in userinput_memory_ids: - mem_item: TextualMemoryItem = mem_cube.text_mem.get(memory_id=memory_id) - mem_type = mem_item.meta_data.memory_type - mem_content = mem_item.memory - - self.log_adding_memory( - memory=mem_content, - memory_type=mem_type, - user_id=msg.user_id, - mem_cube_id=msg.mem_cube_id, - mem_cube=msg.mem_cube, - log_func_callback=self._submit_web_logs, + try: + for user_id in grouped_messages: + for mem_cube_id in grouped_messages[user_id]: + messages = grouped_messages[user_id][mem_cube_id] + if len(messages) == 0: + return + + # for status update + self._set_current_context_from_message(msg=messages[0]) + + # submit logs + for msg in messages: + try: + userinput_memory_ids = json.loads(msg.content) + except Exception as e: + logger.error(f"Error: {e}. Content: {msg.content}", exc_info=True) + userinput_memory_ids = [] + + mem_cube = msg.mem_cube + for memory_id in userinput_memory_ids: + mem_item: TextualMemoryItem = mem_cube.text_mem.get(memory_id=memory_id) + mem_type = mem_item.metadata.memory_type + mem_content = mem_item.memory + + self.log_adding_memory( + memory=mem_content, + memory_type=mem_type, + user_id=msg.user_id, + mem_cube_id=msg.mem_cube_id, + mem_cube=msg.mem_cube, + log_func_callback=self._submit_web_logs, + ) + + # update activation memories + if self.enable_act_memory_update: + self.update_activation_memory_periodically( + interval_seconds=self.monitor.act_mem_update_interval, + label=ADD_LABEL, + user_id=user_id, + mem_cube_id=mem_cube_id, + mem_cube=messages[0].mem_cube, ) - - # update activation memories - if self.enable_act_memory_update: - self.update_activation_memory_periodically( - interval_seconds=self.monitor.act_mem_update_interval, - label=ADD_LABEL, - user_id=user_id, - mem_cube_id=mem_cube_id, - mem_cube=messages[0].mem_cube, - ) + except Exception as e: + logger.error(f"Error: {e}", exc_info=True) def process_session_turn( self, queries: str | list[str], - user_id: str, - mem_cube_id: str, + user_id: UserID | str, + mem_cube_id: MemCubeID | str, mem_cube: GeneralMemCube, top_k: int = 10, ) -> tuple[list[TextualMemoryItem], list[TextualMemoryItem]] | None: @@ -239,7 +292,9 @@ def process_session_turn( results: list[TextualMemoryItem] = self.retriever.search( query=item, mem_cube=mem_cube, top_k=k_per_evidence, method=self.search_method ) - logger.info(f"search results for {missing_evidences}: {results}") + logger.info( + f"search results for {missing_evidences}: {[one.memory for one in results]}" + ) new_candidates.extend(results) if len(new_candidates) == 0: diff --git a/src/memos/mem_scheduler/modules/base.py b/src/memos/mem_scheduler/modules/base.py index 58c645870..37538b0f1 100644 --- a/src/memos/mem_scheduler/modules/base.py +++ b/src/memos/mem_scheduler/modules/base.py @@ -17,8 +17,8 @@ def __init__(self): self._chat_llm = None self._process_llm = None - self._current_mem_cube_id: str | None = None - self._current_mem_cube: GeneralMemCube | None = None + self.current_mem_cube_id: str | None = None + self.current_mem_cube: GeneralMemCube | None = None self.mem_cubes: dict[str, GeneralMemCube] = {} def load_template(self, template_name: str) -> str: @@ -75,9 +75,9 @@ def process_llm(self, value: BaseLLM) -> None: @property def mem_cube(self) -> GeneralMemCube: """The memory cube associated with this MemChat.""" - return self._current_mem_cube + return self.current_mem_cube @mem_cube.setter def mem_cube(self, value: GeneralMemCube) -> None: """The memory cube associated with this MemChat.""" - self._current_mem_cube = value + self.current_mem_cube = value diff --git a/src/memos/mem_scheduler/modules/retriever.py b/src/memos/mem_scheduler/modules/retriever.py index 3430288dd..2cf3bd046 100644 --- a/src/memos/mem_scheduler/modules/retriever.py +++ b/src/memos/mem_scheduler/modules/retriever.py @@ -83,21 +83,22 @@ def rerank_memories( If LLM reranking fails, falls back to original order (truncated to top_k) """ success_flag = False - try: - logger.info(f"Starting memory reranking for {len(original_memories)} memories") - # Build LLM prompt for memory reranking - prompt = self.build_prompt( - "memory_reranking", - queries=[f"[0] {queries[0]}"], - current_order=[f"[{i}] {mem}" for i, mem in enumerate(original_memories)], - ) - logger.debug(f"Generated reranking prompt: {prompt[:200]}...") # Log first 200 chars + logger.info(f"Starting memory reranking for {len(original_memories)} memories") - # Get LLM response - response = self.process_llm.generate([{"role": "user", "content": prompt}]) - logger.debug(f"Received LLM response: {response[:200]}...") # Log first 200 chars + # Build LLM prompt for memory reranking + prompt = self.build_prompt( + "memory_reranking", + queries=[f"[0] {queries[0]}"], + current_order=[f"[{i}] {mem}" for i, mem in enumerate(original_memories)], + ) + logger.debug(f"Generated reranking prompt: {prompt[:200]}...") # Log first 200 chars + # Get LLM response + response = self.process_llm.generate([{"role": "user", "content": prompt}]) + logger.debug(f"Received LLM response: {response[:200]}...") # Log first 200 chars + + try: # Parse JSON response response = extract_json_dict(response) new_order = response["new_order"][:top_k] @@ -109,7 +110,7 @@ def rerank_memories( success_flag = True except Exception as e: logger.error( - f"Failed to rerank memories with LLM;\nException: {e}. ", + f"Failed to rerank memories with LLM. Exception: {e}. Raw response: {response} ", exc_info=True, ) text_memories_with_new_order = original_memories[:top_k] diff --git a/src/memos/mem_scheduler/modules/scheduler_logger.py b/src/memos/mem_scheduler/modules/scheduler_logger.py index d0eecac95..e41b4822e 100644 --- a/src/memos/mem_scheduler/modules/scheduler_logger.py +++ b/src/memos/mem_scheduler/modules/scheduler_logger.py @@ -21,6 +21,7 @@ from memos.mem_scheduler.utils.filter_utils import ( transform_name_to_key, ) +from memos.mem_scheduler.utils.misc_utils import log_exceptions from memos.memories.textual.tree import TextualMemoryItem, TreeTextMemory @@ -34,6 +35,7 @@ def __init__(self): """ super().__init__() + @log_exceptions(logger=logger) def create_autofilled_log_item( self, log_content: str, @@ -47,9 +49,9 @@ def create_autofilled_log_item( text_mem_base: TreeTextMemory = mem_cube.text_mem current_memory_sizes = text_mem_base.get_current_memory_size() current_memory_sizes = { - "long_term_memory_size": current_memory_sizes["LongTermMemory"], - "user_memory_size": current_memory_sizes["UserMemory"], - "working_memory_size": current_memory_sizes["WorkingMemory"], + "long_term_memory_size": current_memory_sizes.get("LongTermMemory", 0), + "user_memory_size": current_memory_sizes.get("UserMemory", 0), + "working_memory_size": current_memory_sizes.get("WorkingMemory", 0), "transformed_act_memory_size": NOT_INITIALIZED, "parameter_memory_size": NOT_INITIALIZED, } @@ -68,8 +70,14 @@ def create_autofilled_log_item( ): activation_monitor = self.monitor.activation_memory_monitors[user_id][mem_cube_id] transformed_act_memory_size = len(activation_monitor.memories) + logger.info( + f'activation_memory_monitors currently has "{transformed_act_memory_size}" transformed memory size' + ) else: transformed_act_memory_size = 0 + logger.info( + f'activation_memory_monitors is not initialized for user "{user_id}" and mem_cube "{mem_cube_id}' + ) current_memory_sizes["transformed_act_memory_size"] = transformed_act_memory_size current_memory_sizes["parameter_memory_size"] = 1 @@ -90,6 +98,7 @@ def create_autofilled_log_item( ) return log_message + @log_exceptions(logger=logger) def log_working_memory_replacement( self, original_memory: list[TextualMemoryItem], @@ -142,6 +151,7 @@ def log_working_memory_replacement( f"transformed to {WORKING_MEMORY_TYPE} memories." ) + @log_exceptions(logger=logger) def log_activation_memory_update( self, original_text_memories: list[str], @@ -185,6 +195,7 @@ def log_activation_memory_update( f"transformed to {WORKING_MEMORY_TYPE} memories." ) + @log_exceptions(logger=logger) def log_adding_memory( self, memory: str, @@ -210,6 +221,7 @@ def log_adding_memory( f"converted to {memory_type} memory in mem_cube {mem_cube_id}: {memory}" ) + @log_exceptions(logger=logger) def validate_schedule_message(self, message: ScheduleMessageItem, label: str): """Validate if the message matches the expected label. @@ -225,6 +237,7 @@ def validate_schedule_message(self, message: ScheduleMessageItem, label: str): return False return True + @log_exceptions(logger=logger) def validate_schedule_messages(self, messages: list[ScheduleMessageItem], label: str): """Validate if all messages match the expected label. diff --git a/src/memos/mem_scheduler/schemas/monitor_schemas.py b/src/memos/mem_scheduler/schemas/monitor_schemas.py index 6cba48fa4..68d53f55b 100644 --- a/src/memos/mem_scheduler/schemas/monitor_schemas.py +++ b/src/memos/mem_scheduler/schemas/monitor_schemas.py @@ -107,7 +107,7 @@ def get_keywords_collections(self) -> Counter: all_keywords = [kw for item in self.queue for kw in item.keywords] return Counter(all_keywords) - def get_queries_with_timesort(self, reverse: bool = True) -> list[dict]: + def get_queries_with_timesort(self, reverse: bool = True) -> list[str]: """ Retrieve all queries sorted by timestamp. diff --git a/src/memos/mem_scheduler/utils/misc_utils.py b/src/memos/mem_scheduler/utils/misc_utils.py index 92b56944c..4a9c0246d 100644 --- a/src/memos/mem_scheduler/utils/misc_utils.py +++ b/src/memos/mem_scheduler/utils/misc_utils.py @@ -1,9 +1,15 @@ import json +from functools import wraps from pathlib import Path import yaml +from memos.log import get_logger + + +logger = get_logger(__name__) + def extract_json_dict(text: str): text = text.strip() @@ -14,8 +20,7 @@ def extract_json_dict(text: str): return res -def parse_yaml(yaml_file): - yaml_path = Path(yaml_file) +def parse_yaml(yaml_file: str | Path): yaml_path = Path(yaml_file) if not yaml_path.is_file(): raise FileNotFoundError(f"No such file: {yaml_file}") @@ -24,3 +29,33 @@ def parse_yaml(yaml_file): data = yaml.safe_load(fr) return data + + +def log_exceptions(logger=logger): + """ + Exception-catching decorator that automatically logs errors (including stack traces) + + Args: + logger: Optional logger object (default: module-level logger) + + Example: + @log_exceptions() + def risky_function(): + raise ValueError("Oops!") + + @log_exceptions(logger=custom_logger) + def another_risky_function(): + might_fail() + """ + + def decorator(func): + @wraps(func) + def wrapper(*args, **kwargs): + try: + return func(*args, **kwargs) + except Exception as e: + logger.error(f"Error in {func.__name__}: {e}", exc_info=True) + + return wrapper + + return decorator diff --git a/src/memos/mem_user/factory.py b/src/memos/mem_user/factory.py new file mode 100644 index 000000000..060723c83 --- /dev/null +++ b/src/memos/mem_user/factory.py @@ -0,0 +1,94 @@ +from typing import Any, ClassVar + +from memos.configs.mem_user import UserManagerConfigFactory +from memos.mem_user.mysql_user_manager import MySQLUserManager +from memos.mem_user.user_manager import UserManager + + +class UserManagerFactory: + """Factory class for creating user manager instances.""" + + backend_to_class: ClassVar[dict[str, Any]] = { + "sqlite": UserManager, + "mysql": MySQLUserManager, + } + + @classmethod + def from_config( + cls, config_factory: UserManagerConfigFactory + ) -> UserManager | MySQLUserManager: + """Create a user manager instance from configuration. + + Args: + config_factory: Configuration factory containing backend and config + + Returns: + User manager instance + + Raises: + ValueError: If backend is not supported + """ + backend = config_factory.backend + if backend not in cls.backend_to_class: + raise ValueError(f"Invalid user manager backend: {backend}") + + user_manager_class = cls.backend_to_class[backend] + config = config_factory.config + + # Use model_dump() to convert Pydantic model to dict and unpack as kwargs + return user_manager_class(**config.model_dump()) + + @classmethod + def create_sqlite(cls, db_path: str | None = None, user_id: str = "root") -> UserManager: + """Create SQLite user manager with default configuration. + + Args: + db_path: Path to SQLite database file + user_id: Default user ID for initialization + + Returns: + SQLite user manager instance + """ + config_factory = UserManagerConfigFactory( + backend="sqlite", config={"db_path": db_path, "user_id": user_id} + ) + return cls.from_config(config_factory) + + @classmethod + def create_mysql( + cls, + user_id: str = "root", + host: str = "localhost", + port: int = 3306, + username: str = "root", + password: str = "", + database: str = "memos_users", + charset: str = "utf8mb4", + ) -> MySQLUserManager: + """Create MySQL user manager with specified configuration. + + Args: + user_id: Default user ID for initialization + host: MySQL server host + port: MySQL server port + username: MySQL username + password: MySQL password + database: MySQL database name + charset: MySQL charset + + Returns: + MySQL user manager instance + """ + config_factory = UserManagerConfigFactory( + backend="mysql", + config={ + "user_id": user_id, + "host": host, + "port": port, + "username": username, + "password": password, + "database": database, + "charset": charset, + }, + ) + return cls.from_config(config_factory) diff --git a/src/memos/mem_user/mysql_persistent_user_manager.py b/src/memos/mem_user/mysql_persistent_user_manager.py new file mode 100644 index 000000000..60a4a1032 --- /dev/null +++ b/src/memos/mem_user/mysql_persistent_user_manager.py @@ -0,0 +1,271 @@ +"""Persistent user management system for MemOS with configuration storage. + +This module extends the MySQL UserManager to provide persistent storage +for user configurations and MOS instances. +""" + +import json + +from datetime import datetime +from typing import Any + +from sqlalchemy import Column, String, Text + +from memos.configs.mem_os import MOSConfig +from memos.log import get_logger +from memos.mem_user.mysql_user_manager import Base, MySQLUserManager + + +logger = get_logger(__name__) + + +class UserConfig(Base): + """User configuration model for the database.""" + + __tablename__ = "user_configs" + + user_id = Column(String, primary_key=True) + config_data = Column(Text, nullable=False) # JSON string of MOSConfig + created_at = Column(String, nullable=False) # ISO format timestamp + updated_at = Column(String, nullable=False) # ISO format timestamp + + def __repr__(self): + return f"" + + +class MySQLPersistentUserManager(MySQLUserManager): + """Extended MySQLUserManager with configuration persistence.""" + + def __init__( + self, + user_id: str = "root", + host: str = "localhost", + port: int = 3306, + username: str = "root", + password: str = "", + database: str = "memos_users", + charset: str = "utf8mb4", + ): + """Initialize the persistent user manager. + + Args: + user_id (str, optional): User ID. If None, uses default user ID. + host (str): MySQL server host. Defaults to "localhost". + port (int): MySQL server port. Defaults to 3306. + username (str): MySQL username. Defaults to "root". + password (str): MySQL password. Defaults to "". + database (str): MySQL database name. Defaults to "memos_users". + charset (str): MySQL charset. Defaults to "utf8mb4". + """ + super().__init__(user_id, host, port, username, password, database, charset) + + # Create user_configs table + Base.metadata.create_all(bind=self.engine) + logger.info("MySQLPersistentUserManager initialized with configuration storage") + + def _convert_datetime_strings(self, obj: Any) -> Any: + """Recursively convert datetime strings back to datetime objects in config dict. + + Args: + obj: The object to process (dict, list, or primitive type) + + Returns: + The object with datetime strings converted to datetime objects + """ + if isinstance(obj, dict): + result = {} + for key, value in obj.items(): + if key == "created_at" and isinstance(value, str): + try: + result[key] = datetime.fromisoformat(value) + except ValueError: + # If parsing fails, keep the original string + result[key] = value + else: + result[key] = self._convert_datetime_strings(value) + return result + elif isinstance(obj, list): + return [self._convert_datetime_strings(item) for item in obj] + else: + return obj + + def save_user_config(self, user_id: str, config: MOSConfig) -> bool: + """Save user configuration to database. + + Args: + user_id (str): The user ID. + config (MOSConfig): The user's MOS configuration. + + Returns: + bool: True if successful, False otherwise. + """ + session = self._get_session() + try: + # Convert config to JSON string with proper datetime handling + config_dict = config.model_dump(mode="json") + config_json = json.dumps(config_dict, indent=2) + + now = datetime.now().isoformat() + + # Check if config already exists + existing_config = ( + session.query(UserConfig).filter(UserConfig.user_id == user_id).first() + ) + + if existing_config: + # Update existing config + existing_config.config_data = config_json + existing_config.updated_at = now + logger.info(f"Updated configuration for user {user_id}") + else: + # Create new config + user_config = UserConfig( + user_id=user_id, config_data=config_json, created_at=now, updated_at=now + ) + session.add(user_config) + logger.info(f"Saved new configuration for user {user_id}") + + session.commit() + return True + + except Exception as e: + session.rollback() + logger.error(f"Error saving user config for {user_id}: {e}") + return False + finally: + session.close() + + def get_user_config(self, user_id: str) -> MOSConfig | None: + """Get user configuration from database. + + Args: + user_id (str): The user ID. + + Returns: + MOSConfig | None: The user's configuration or None if not found. + """ + session = self._get_session() + try: + user_config = session.query(UserConfig).filter(UserConfig.user_id == user_id).first() + + if user_config: + config_dict = json.loads(user_config.config_data) + # Convert datetime strings back to datetime objects + config_dict = self._convert_datetime_strings(config_dict) + return MOSConfig(**config_dict) + return None + + except Exception as e: + logger.error(f"Error loading user config for {user_id}: {e}") + return None + finally: + session.close() + + def delete_user_config(self, user_id: str) -> bool: + """Delete user configuration from database. + + Args: + user_id (str): The user ID. + + Returns: + bool: True if successful, False otherwise. + """ + session = self._get_session() + try: + user_config = session.query(UserConfig).filter(UserConfig.user_id == user_id).first() + + if user_config: + session.delete(user_config) + session.commit() + logger.info(f"Deleted configuration for user {user_id}") + return True + return False + + except Exception as e: + session.rollback() + logger.error(f"Error deleting user config for {user_id}: {e}") + return False + finally: + session.close() + + def list_user_configs(self) -> dict[str, MOSConfig]: + """List all user configurations. + + Returns: + Dict[str, MOSConfig]: Dictionary mapping user_id to MOSConfig. + """ + session = self._get_session() + try: + user_configs = session.query(UserConfig).all() + result = {} + + for user_config in user_configs: + try: + config_dict = json.loads(user_config.config_data) + # Convert datetime strings back to datetime objects + config_dict = self._convert_datetime_strings(config_dict) + result[user_config.user_id] = MOSConfig(**config_dict) + except Exception as e: + logger.error(f"Error parsing config for user {user_config.user_id}: {e}") + continue + + return result + + except Exception as e: + logger.error(f"Error listing user configs: {e}") + return {} + finally: + session.close() + + def create_user_with_config( + self, user_name: str, config: MOSConfig, role=None, user_id: str | None = None + ) -> str: + """Create a new user with configuration. + + Args: + user_name (str): Name of the user. + config (MOSConfig): The user's configuration. + role: User role (optional, uses default from UserManager). + user_id (str, optional): Custom user ID. + + Returns: + str: The created user ID. + + Raises: + ValueError: If user_name already exists. + """ + # Create user using parent method + created_user_id = self.create_user(user_name, role, user_id) + + # Save configuration + if not self.save_user_config(created_user_id, config): + logger.error(f"Failed to save configuration for user {created_user_id}") + + return created_user_id + + def delete_user(self, user_id: str) -> bool: + """Delete a user and their configuration. + + Args: + user_id (str): The user ID. + + Returns: + bool: True if successful, False otherwise. + """ + # Delete configuration first + self.delete_user_config(user_id) + + # Delete user using parent method + return super().delete_user(user_id) + + def get_user_cube_access(self, user_id: str) -> list[str]: + """Get list of cube IDs that a user has access to. + + Args: + user_id (str): The user ID. + + Returns: + list[str]: List of cube IDs the user can access. + """ + cubes = self.get_user_cubes(user_id) + return [cube.cube_id for cube in cubes] diff --git a/src/memos/mem_user/mysql_user_manager.py b/src/memos/mem_user/mysql_user_manager.py new file mode 100644 index 000000000..d7f77c024 --- /dev/null +++ b/src/memos/mem_user/mysql_user_manager.py @@ -0,0 +1,503 @@ +"""User management system for MemOS. + +This module provides user authentication, authorization, and cube management +functionality using SQLAlchemy and MySQL. +""" + +import uuid + +from datetime import datetime +from enum import Enum + +from sqlalchemy import ( + Boolean, + Column, + DateTime, + ForeignKey, + String, + Table, + create_engine, +) +from sqlalchemy import ( + Enum as SQLEnum, +) +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session, declarative_base, relationship, sessionmaker + +from memos.log import get_logger + + +logger = get_logger(__name__) + +Base = declarative_base() + + +class UserRole(Enum): + """User roles enumeration.""" + + ROOT = "root" + ADMIN = "admin" + USER = "user" + GUEST = "guest" + + +# Association table for many-to-many relationship between users and cubes +user_cube_association = Table( + "user_cube_association", + Base.metadata, + Column("user_id", String, ForeignKey("users.user_id"), primary_key=True), + Column("cube_id", String, ForeignKey("cubes.cube_id"), primary_key=True), + Column("created_at", DateTime, default=datetime.now), +) + + +class User(Base): + """User model for the database.""" + + __tablename__ = "users" + + user_id = Column(String, primary_key=True, default=lambda: str(uuid.uuid4())) + user_name = Column(String, unique=True, nullable=False) + role = Column(SQLEnum(UserRole), default=UserRole.USER, nullable=False) + created_at = Column(DateTime, default=datetime.now, nullable=False) + updated_at = Column(DateTime, default=datetime.now, onupdate=datetime.now, nullable=False) + is_active = Column(Boolean, default=True, nullable=False) + + # Relationship with cubes + cubes = relationship("Cube", secondary=user_cube_association, back_populates="users") + owned_cubes = relationship("Cube", back_populates="owner", cascade="all, delete-orphan") + + def __repr__(self): + return f"" + + +class Cube(Base): + """Cube model for the database.""" + + __tablename__ = "cubes" + + cube_id = Column(String, primary_key=True, default=lambda: str(uuid.uuid4())) + cube_name = Column(String, nullable=False) + cube_path = Column(String, nullable=True) # Local path or remote repo + owner_id = Column(String, ForeignKey("users.user_id"), nullable=False) + created_at = Column(DateTime, default=datetime.now, nullable=False) + updated_at = Column(DateTime, default=datetime.now, onupdate=datetime.now, nullable=False) + is_active = Column(Boolean, default=True, nullable=False) + + # Relationships + owner = relationship("User", back_populates="owned_cubes") + users = relationship("User", secondary=user_cube_association, back_populates="cubes") + + def __repr__(self): + return f"" + + +class MySQLUserManager: + """User management system for MemOS using MySQL.""" + + def __init__( + self, + user_id: str = "root", + host: str = "localhost", + port: int = 3306, + username: str = "root", + password: str = "", + database: str = "memos_users", + charset: str = "utf8mb4", + ): + """Initialize the user manager with MySQL database connection. + + Args: + user_id (str, optional): User ID. If None, uses default user ID. + host (str): MySQL server host. Defaults to "localhost". + port (int): MySQL server port. Defaults to 3306. + username (str): MySQL username. Defaults to "root". + password (str): MySQL password. Defaults to "". + database (str): MySQL database name. Defaults to "memos_users". + charset (str): MySQL charset. Defaults to "utf8mb4". + """ + # Build MySQL connection URL + if password: + connection_url = ( + f"mysql+pymysql://{username}:{password}@{host}:{port}/{database}?charset={charset}" + ) + else: + connection_url = ( + f"mysql+pymysql://{username}@{host}:{port}/{database}?charset={charset}" + ) + + self.connection_url = connection_url + self.engine = create_engine(connection_url, echo=False, pool_pre_ping=True) + self.SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=self.engine) + + # Create tables + Base.metadata.create_all(bind=self.engine) + + # Initialize with root user if no users exist + self._init_root_user(user_id) + + logger.info(f"MySQLUserManager initialized with database at {host}:{port}/{database}") + + def _get_session(self) -> Session: + """Get a database session.""" + return self.SessionLocal() + + def _init_root_user(self, user_id: str) -> None: + """Initialize the root user if no users exist.""" + session = self._get_session() + try: + # Check if any users exist + user_count = session.query(User).count() + if user_count == 0: + root_user = User(user_id=user_id, user_name=user_id, role=UserRole.ROOT) + session.add(root_user) + session.commit() + logger.info("Root user created successfully") + else: + self.create_user(user_name=user_id, user_id=user_id, role=UserRole.ROOT) + except Exception as e: + session.rollback() + logger.error(f"Failed to create {user_id} user: {e}") + finally: + session.close() + + def create_user( + self, user_name: str, role: UserRole = UserRole.USER, user_id: str | None = None + ) -> str: + """Create a new user. + + Args: + user_name (str): Name of the user. + role (UserRole): Role of the user. + user_id (str, optional): Custom user ID. If None, generates UUID. + + Returns: + str: The created user ID. + + Raises: + ValueError: If user_name already exists. + """ + session = self._get_session() + try: + # Check if user_name already exists + existing_user = session.query(User).filter(User.user_name == user_name).first() + if existing_user: + logger.info(f"User with name '{user_name}' already exists") + return existing_user.user_id + user = User(user_name=user_name, role=role, user_id=user_id or str(uuid.uuid4())) + session.add(user) + session.commit() + logger.info(f"User '{user_name}' created with ID: {user.user_id}") + return user.user_id + except IntegrityError: + session.rollback() + logger.info(f"Failed to create user with name '{user_name}' already exists") + except Exception as e: + session.rollback() + logger.error(f"Error creating user: {e}") + raise + finally: + session.close() + + def get_user(self, user_id: str) -> User | None: + """Get user by ID. + + Args: + user_id (str): The user ID. + + Returns: + User: The user object or None if not found. + """ + session = self._get_session() + try: + return session.query(User).filter(User.user_id == user_id).first() + finally: + session.close() + + def get_user_by_name(self, user_name: str) -> User | None: + """Get user by name. + + Args: + user_name (str): The user name. + + Returns: + User: The user object or None if not found. + """ + session = self._get_session() + try: + return session.query(User).filter(User.user_name == user_name).first() + finally: + session.close() + + def validate_user(self, user_id: str) -> bool: + """Validate if a user exists and is active. + + Args: + user_id (str): The user ID to validate. + + Returns: + bool: True if user exists and is active, False otherwise. + """ + user = self.get_user(user_id) + return user is not None and user.is_active + + def list_users(self) -> list[User]: + """List all active users. + + Returns: + list[User]: List of all active users. + """ + session = self._get_session() + try: + return session.query(User).filter(User.is_active).all() + finally: + session.close() + + def create_cube( + self, + cube_name: str, + owner_id: str, + cube_path: str | None = None, + cube_id: str | None = None, + ) -> str: + """Create a new cube. + + Args: + cube_name (str): Name of the cube. + owner_id (str): ID of the cube owner. + cube_path (str, optional): Path to the cube. + cube_id (str, optional): Custom cube ID. If None, generates UUID. + + Returns: + str: The created cube ID. + + Raises: + ValueError: If owner doesn't exist. + """ + session = self._get_session() + try: + # Validate owner exists + owner = session.query(User).filter(User.user_id == owner_id).first() + if not owner: + raise ValueError(f"User with ID '{owner_id}' does not exist") + + cube = Cube( + cube_name=cube_name, + owner_id=owner_id, + cube_path=cube_path, + cube_id=cube_id or str(uuid.uuid4()), + ) + session.add(cube) + + # Add owner to cube users + cube.users.append(owner) + + session.commit() + logger.info(f"Cube '{cube_name}' created with ID: {cube.cube_id}") + return cube.cube_id + except Exception as e: + session.rollback() + logger.error(f"Error creating cube: {e}") + raise + finally: + session.close() + + def get_cube(self, cube_id: str) -> Cube | None: + """Get cube by ID. + + Args: + cube_id (str): The cube ID. + + Returns: + Cube: The cube object or None if not found. + """ + session = self._get_session() + try: + return session.query(Cube).filter(Cube.cube_id == cube_id).first() + finally: + session.close() + + def validate_user_cube_access(self, user_id: str, cube_id: str) -> bool: + """Validate if a user has access to a cube. + + Args: + user_id (str): The user ID. + cube_id (str): The cube ID. + + Returns: + bool: True if user has access to cube, False otherwise. + """ + session = self._get_session() + try: + # Check if user exists and is active + user = session.query(User).filter(User.user_id == user_id, User.is_active).first() + if not user: + return False + + # Check if cube exists and is active + cube = session.query(Cube).filter(Cube.cube_id == cube_id, Cube.is_active).first() + if not cube: + return False + + # Check if user has access to cube (owner or in users list) + if cube.owner_id == user_id: + return True + + # Check many-to-many relationship + return user in cube.users + finally: + session.close() + + def get_user_cubes(self, user_id: str) -> list[Cube]: + """Get all cubes accessible by a user. + + Args: + user_id (str): The user ID. + + Returns: + list[Cube]: List of cubes accessible by the user. + """ + session = self._get_session() + try: + user = session.query(User).filter(User.user_id == user_id).first() + if not user: + return [] + + active_cubes = [cube for cube in user.cubes if cube.is_active] + return sorted(active_cubes, key=lambda cube: cube.created_at, reverse=True) + finally: + session.close() + + def add_user_to_cube(self, user_id: str, cube_id: str) -> bool: + """Add a user to a cube's access list. + + Args: + user_id (str): The user ID. + cube_id (str): The cube ID. + + Returns: + bool: True if successful, False otherwise. + """ + session = self._get_session() + try: + user = session.query(User).filter(User.user_id == user_id).first() + cube = session.query(Cube).filter(Cube.cube_id == cube_id).first() + + if not user or not cube: + return False + + if user not in cube.users: + cube.users.append(user) + session.commit() + logger.info(f"User '{user_id}' added to cube '{cube_id}'") + + return True + except Exception as e: + session.rollback() + logger.error(f"Error adding user to cube: {e}") + return False + finally: + session.close() + + def remove_user_from_cube(self, user_id: str, cube_id: str) -> bool: + """Remove a user from a cube's access list. + + Args: + user_id (str): The user ID. + cube_id (str): The cube ID. + + Returns: + bool: True if successful, False otherwise. + """ + session = self._get_session() + try: + user = session.query(User).filter(User.user_id == user_id).first() + cube = session.query(Cube).filter(Cube.cube_id == cube_id).first() + + if not user or not cube: + return False + + # Don't remove owner + if cube.owner_id == user_id: + logger.warning(f"Cannot remove owner '{user_id}' from cube '{cube_id}'") + return False + + if user in cube.users: + cube.users.remove(user) + session.commit() + logger.info(f"User '{user_id}' removed from cube '{cube_id}'") + + return True + except Exception as e: + session.rollback() + logger.error(f"Error removing user from cube: {e}") + return False + finally: + session.close() + + def delete_user(self, user_id: str) -> bool: + """Soft delete a user (set is_active to False). + + Args: + user_id (str): The user ID. + + Returns: + bool: True if successful, False otherwise. + """ + session = self._get_session() + try: + user = session.query(User).filter(User.user_id == user_id).first() + if not user: + return False + + # Don't delete root user + if user.role == UserRole.ROOT: + logger.warning("Cannot delete root user") + return False + + user.is_active = False + session.commit() + logger.info(f"User '{user_id}' deactivated") + return True + except Exception as e: + session.rollback() + logger.error(f"Error deleting user: {e}") + return False + finally: + session.close() + + def delete_cube(self, cube_id: str) -> bool: + """Soft delete a cube (set is_active to False). + + Args: + cube_id (str): The cube ID. + + Returns: + bool: True if successful, False otherwise. + """ + session = self._get_session() + try: + cube = session.query(Cube).filter(Cube.cube_id == cube_id).first() + if not cube: + return False + + cube.is_active = False + session.commit() + logger.info(f"Cube '{cube_id}' deactivated") + return True + except Exception as e: + session.rollback() + logger.error(f"Error deleting cube: {e}") + return False + finally: + session.close() + + def close(self) -> None: + """Close the database engine and dispose of all connections. + + This method should be called when the MySQLUserManager is no longer needed + to ensure proper cleanup of database connections. + """ + if hasattr(self, "engine"): + self.engine.dispose() + logger.info("MySQLUserManager database connections closed") diff --git a/src/memos/mem_user/persistent_factory.py b/src/memos/mem_user/persistent_factory.py new file mode 100644 index 000000000..b5ece61b5 --- /dev/null +++ b/src/memos/mem_user/persistent_factory.py @@ -0,0 +1,96 @@ +from typing import Any, ClassVar + +from memos.configs.mem_user import UserManagerConfigFactory +from memos.mem_user.mysql_persistent_user_manager import MySQLPersistentUserManager +from memos.mem_user.persistent_user_manager import PersistentUserManager + + +class PersistentUserManagerFactory: + """Factory class for creating persistent user manager instances.""" + + backend_to_class: ClassVar[dict[str, Any]] = { + "sqlite": PersistentUserManager, + "mysql": MySQLPersistentUserManager, + } + + @classmethod + def from_config( + cls, config_factory: UserManagerConfigFactory + ) -> PersistentUserManager | MySQLPersistentUserManager: + """Create a persistent user manager instance from configuration. + + Args: + config_factory: Configuration factory containing backend and config + + Returns: + Persistent user manager instance + + Raises: + ValueError: If backend is not supported + """ + backend = config_factory.backend + if backend not in cls.backend_to_class: + raise ValueError(f"Invalid persistent user manager backend: {backend}") + + user_manager_class = cls.backend_to_class[backend] + config = config_factory.config + + # Use model_dump() to convert Pydantic model to dict and unpack as kwargs + return user_manager_class(**config.model_dump()) + + @classmethod + def create_sqlite( + cls, db_path: str | None = None, user_id: str = "root" + ) -> PersistentUserManager: + """Create SQLite persistent user manager with default configuration. + + Args: + db_path: Path to SQLite database file + user_id: Default user ID for initialization + + Returns: + SQLite persistent user manager instance + """ + config_factory = UserManagerConfigFactory( + backend="sqlite", config={"db_path": db_path, "user_id": user_id} + ) + return cls.from_config(config_factory) + + @classmethod + def create_mysql( + cls, + user_id: str = "root", + host: str = "localhost", + port: int = 3306, + username: str = "root", + password: str = "", + database: str = "memos_users", + charset: str = "utf8mb4", + ) -> MySQLPersistentUserManager: + """Create MySQL persistent user manager with specified configuration. + + Args: + user_id: Default user ID for initialization + host: MySQL server host + port: MySQL server port + username: MySQL username + password: MySQL password + database: MySQL database name + charset: MySQL charset + + Returns: + MySQL persistent user manager instance + """ + config_factory = UserManagerConfigFactory( + backend="mysql", + config={ + "user_id": user_id, + "host": host, + "port": port, + "username": username, + "password": password, + "database": database, + "charset": charset, + }, + ) + return cls.from_config(config_factory) diff --git a/src/memos/memories/textual/general.py b/src/memos/memories/textual/general.py index 2754fc1ae..4a1d90cb0 100644 --- a/src/memos/memories/textual/general.py +++ b/src/memos/memories/textual/general.py @@ -17,6 +17,7 @@ from memos.vec_dbs.factory import QdrantVecDB, VecDBFactory from memos.vec_dbs.item import VecDBItem + logger = get_logger(__name__) @@ -36,11 +37,7 @@ def __init__(self, config: GeneralTextMemoryConfig): stop=stop_after_attempt(3), retry=retry_if_exception_type(json.JSONDecodeError), before_sleep=lambda retry_state: logger.warning( - "Extracting memory failed due to JSON decode error: {error}, Attempt retry: {attempt_number} / {max_attempt_number}".format( - error=retry_state.outcome.exception(), - attempt_number=retry_state.attempt_number, - max_attempt_number=3, - ) + f"Extracting memory failed due to JSON decode error: {retry_state.outcome.exception()}, Attempt retry: {retry_state.attempt_number} / {3}" ), ) def extract(self, messages: MessageList) -> list[TextualMemoryItem]: diff --git a/src/memos/memories/textual/item.py b/src/memos/memories/textual/item.py index 06d832b39..c287c1918 100644 --- a/src/memos/memories/textual/item.py +++ b/src/memos/memories/textual/item.py @@ -59,7 +59,7 @@ def __str__(self) -> str: class TreeNodeTextualMemoryMetadata(TextualMemoryMetadata): """Extended metadata for structured memory, layered retrieval, and lifecycle tracking.""" - memory_type: Literal["WorkingMemory", "LongTermMemory", "UserMemory"] = Field( + memory_type: Literal["WorkingMemory", "LongTermMemory", "UserMemory", "OuterMemory"] = Field( default="WorkingMemory", description="Memory lifecycle type." ) sources: list[str] | None = Field( diff --git a/src/memos/memories/textual/tree.py b/src/memos/memories/textual/tree.py index 3b50001d8..601597b19 100644 --- a/src/memos/memories/textual/tree.py +++ b/src/memos/memories/textual/tree.py @@ -117,13 +117,19 @@ def search( logger.warning( "Internet retriever is init by config , but this search set manual_close_internet is True and will close it" ) - self.internet_retriever = None - searcher = Searcher( - self.dispatcher_llm, - self.graph_store, - self.embedder, - internet_retriever=self.internet_retriever, - ) + searcher = Searcher( + self.dispatcher_llm, + self.graph_store, + self.embedder, + internet_retriever=None, + ) + else: + searcher = Searcher( + self.dispatcher_llm, + self.graph_store, + self.embedder, + internet_retriever=self.internet_retriever, + ) return searcher.search(query, top_k, info, mode, memory_type) def get_relevant_subgraph( diff --git a/src/memos/memories/textual/tree_text_memory/organize/conflict.py b/src/memos/memories/textual/tree_text_memory/organize/conflict.py index ccb86e321..2ea16ed2f 100644 --- a/src/memos/memories/textual/tree_text_memory/organize/conflict.py +++ b/src/memos/memories/textual/tree_text_memory/organize/conflict.py @@ -2,7 +2,9 @@ import re from datetime import datetime + from dateutil import parser + from memos.embedders.base import BaseEmbedder from memos.graph_dbs.neo4j import Neo4jGraphDB from memos.llms.base import BaseLLM diff --git a/src/memos/memories/textual/tree_text_memory/organize/relation_reason_detector.py b/src/memos/memories/textual/tree_text_memory/organize/relation_reason_detector.py index cc755d6dd..4fca0be83 100644 --- a/src/memos/memories/textual/tree_text_memory/organize/relation_reason_detector.py +++ b/src/memos/memories/textual/tree_text_memory/organize/relation_reason_detector.py @@ -1,4 +1,5 @@ import json +import traceback from memos.embedders.factory import OllamaEmbedder from memos.graph_dbs.item import GraphDBNode @@ -30,53 +31,57 @@ def process_node(self, node: GraphDBNode, exclude_ids: list[str], top_k: int = 5 3) Sequence links 4) Aggregate concepts """ - if node.metadata.type == "reasoning": - logger.info(f"Skip reasoning for inferred node {node.id}") - return { - "relations": [], - "inferred_nodes": [], - "sequence_links": [], - "aggregate_nodes": [], - } - results = { "relations": [], "inferred_nodes": [], "sequence_links": [], "aggregate_nodes": [], } + try: + if node.metadata.type == "reasoning": + logger.info(f"Skip reasoning for inferred node {node.id}") + return { + "relations": [], + "inferred_nodes": [], + "sequence_links": [], + "aggregate_nodes": [], + } + + nearest = self.graph_store.get_neighbors_by_tag( + tags=node.metadata.tags, + exclude_ids=exclude_ids, + top_k=top_k, + min_overlap=2, + ) + nearest = [GraphDBNode(**cand_data) for cand_data in nearest] + + """ + # 1) Pairwise relations (including CAUSE/CONDITION/CONFLICT) + pairwise = self._detect_pairwise_causal_condition_relations(node, nearest) + results["relations"].extend(pairwise["relations"]) + """ + + """ + # 2) Inferred nodes (from causal/condition) + inferred = self._infer_fact_nodes_from_relations(pairwise) + results["inferred_nodes"].extend(inferred) + """ + + """ + 3) Sequence (optional, if you have timestamps) + seq = self._detect_sequence_links(node, nearest) + results["sequence_links"].extend(seq) + """ + + # 4) Aggregate + agg = self._detect_aggregate_node_for_group(node, nearest, min_group_size=5) + if agg: + results["aggregate_nodes"].append(agg) - nearest = self.graph_store.get_neighbors_by_tag( - tags=node.metadata.tags, - exclude_ids=exclude_ids, - top_k=top_k, - min_overlap=2, - ) - nearest = [GraphDBNode(**cand_data) for cand_data in nearest] - - """ - # 1) Pairwise relations (including CAUSE/CONDITION/CONFLICT) - pairwise = self._detect_pairwise_causal_condition_relations(node, nearest) - results["relations"].extend(pairwise["relations"]) - """ - - """ - # 2) Inferred nodes (from causal/condition) - inferred = self._infer_fact_nodes_from_relations(pairwise) - results["inferred_nodes"].extend(inferred) - """ - - """ - 3) Sequence (optional, if you have timestamps) - seq = self._detect_sequence_links(node, nearest) - results["sequence_links"].extend(seq) - """ - - # 4) Aggregate - agg = self._detect_aggregate_node_for_group(node, nearest, min_group_size=5) - if agg: - results["aggregate_nodes"].append(agg) - + except Exception as e: + logger.error( + f"Error {e} while process struct reorganize: trace: {traceback.format_exc()}" + ) return results def _detect_pairwise_causal_condition_relations( @@ -176,10 +181,9 @@ def _detect_aggregate_node_for_group( joined = "\n".join(f"- {n.memory}" for n in combined_nodes) prompt = AGGREGATE_PROMPT.replace("{joined}", joined) response_text = self._call_llm(prompt) - response_json = self._parse_json_result(response_text) - if not response_json: + summary = self._parse_json_result(response_text) + if not summary: return None - summary = json.loads(response_text) embedding = self.embedder.embed([summary["value"]])[0] parent_node = GraphDBNode( diff --git a/src/memos/memories/textual/tree_text_memory/organize/reorganizer.py b/src/memos/memories/textual/tree_text_memory/organize/reorganizer.py index 43b0217be..e593857e5 100644 --- a/src/memos/memories/textual/tree_text_memory/organize/reorganizer.py +++ b/src/memos/memories/textual/tree_text_memory/organize/reorganizer.py @@ -125,8 +125,8 @@ def _run_structure_organizer_loop(self): """ import schedule - schedule.every(20).seconds.do(self.optimize_structure, scope="LongTermMemory") - schedule.every(20).seconds.do(self.optimize_structure, scope="UserMemory") + schedule.every(600).seconds.do(self.optimize_structure, scope="LongTermMemory") + schedule.every(600).seconds.do(self.optimize_structure, scope="UserMemory") logger.info("Structure optimizer schedule started.") while not getattr(self, "_stop_scheduler", False): @@ -198,7 +198,7 @@ def optimize_structure( logger.info(f"Already optimizing for {scope}. Skipping.") return - if self.graph_store.count_nodes(scope) == 0: + if self.graph_store.node_not_exist(scope): logger.debug(f"[GraphStructureReorganize] No nodes for scope={scope}. Skip.") return @@ -251,7 +251,10 @@ def optimize_structure( try: f.result() except Exception as e: - logger.warning(f"[Reorganize] Cluster processing failed: {e}") + logger.warning( + f"[Reorganize] Cluster processing " + f"failed: {e}, trace: {traceback.format_exc()}" + ) logger.info("[GraphStructure Reorganize] Structure optimization finished.") finally: @@ -343,7 +346,7 @@ def _process_cluster_and_write( agg_node.metadata.model_dump(exclude_none=True), ) for child_id in agg_node.metadata.sources: - self.graph_store.add_edge(agg_node.id, child_id, "AGGREGATES") + self.graph_store.add_edge(agg_node.id, child_id, "AGGREGATE_TO") logger.info("[Reorganizer] Cluster relation/reasoning done.") diff --git a/src/memos/memories/textual/tree_text_memory/retrieve/internet_retriever.py b/src/memos/memories/textual/tree_text_memory/retrieve/internet_retriever.py index de4f3646b..8892a7232 100644 --- a/src/memos/memories/textual/tree_text_memory/retrieve/internet_retriever.py +++ b/src/memos/memories/textual/tree_text_memory/retrieve/internet_retriever.py @@ -127,7 +127,7 @@ def __init__( self.embedder = embedder def retrieve_from_internet( - self, query: str, top_k: int = 10, parsed_goal=None + self, query: str, top_k: int = 10, parsed_goal=None, info=None ) -> list[TextualMemoryItem]: """ Retrieve information from the internet and convert to TextualMemoryItem format @@ -136,6 +136,7 @@ def retrieve_from_internet( query: Search query top_k: Number of results to return parsed_goal: Parsed task goal (optional) + info (dict): Leave a record of memory consumption. Returns: List of TextualMemoryItem @@ -157,8 +158,8 @@ def retrieve_from_internet( memory_content = f"Title: {title}\nSummary: {snippet}\nSource: {link}" # Create metadata metadata = TreeNodeTextualMemoryMetadata( - user_id=None, - session_id=None, + user_id=info.get("user_id", ""), + session_id=info.get("session_id", ""), status="activated", type="fact", # Internet search results are usually factual information memory_time=datetime.now().strftime("%Y-%m-%d"), diff --git a/src/memos/memories/textual/tree_text_memory/retrieve/internet_retriever_factory.py b/src/memos/memories/textual/tree_text_memory/retrieve/internet_retriever_factory.py index d6af5944c..39414c3eb 100644 --- a/src/memos/memories/textual/tree_text_memory/retrieve/internet_retriever_factory.py +++ b/src/memos/memories/textual/tree_text_memory/retrieve/internet_retriever_factory.py @@ -4,6 +4,7 @@ from memos.configs.internet_retriever import InternetRetrieverConfigFactory from memos.embedders.base import BaseEmbedder +from memos.mem_reader.factory import MemReaderFactory from memos.memories.textual.tree_text_memory.retrieve.internet_retriever import ( InternetGoogleRetriever, ) @@ -66,6 +67,7 @@ def from_config( access_key=config.api_key, # Use api_key as access_key for xinyu search_engine_id=config.search_engine_id, embedder=embedder, + reader=MemReaderFactory.from_config(config.reader), max_results=config.max_results, ) else: diff --git a/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py b/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py index 40bd01a4d..6cfd37903 100644 --- a/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py +++ b/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py @@ -136,12 +136,12 @@ def retrieve_from_internet(): """ Retrieve information from the internet using Google Custom Search API. """ - if not self.internet_retriever: + if not self.internet_retriever or mode == "fast": return [] if memory_type not in ["All"]: return [] internet_items = self.internet_retriever.retrieve_from_internet( - query=query, top_k=top_k, parsed_goal=parsed_goal + query=query, top_k=top_k, parsed_goal=parsed_goal, info=info ) # Convert to the format expected by reranker @@ -149,7 +149,7 @@ def retrieve_from_internet(): query=query, query_embedding=query_embedding[0], graph_results=internet_items, - top_k=top_k * 2, + top_k=min(top_k, 5), parsed_goal=parsed_goal, ) return ranked_memories @@ -184,14 +184,6 @@ def retrieve_from_internet(): TextualMemoryItem(id=item.id, memory=item.memory, metadata=new_meta) ) - # Step 4: Reasoning over all retrieved and ranked memory - if mode == "fine": - searched_res = self.reasoner.reason( - query=query, - ranked_memories=searched_res, - parsed_goal=parsed_goal, - ) - # Step 5: Update usage history with current timestamp now_time = datetime.now().isoformat() usage_record = json.dumps( diff --git a/src/memos/memories/textual/tree_text_memory/retrieve/xinyusearch.py b/src/memos/memories/textual/tree_text_memory/retrieve/xinyusearch.py index b803dfa4c..6e1b8ea18 100644 --- a/src/memos/memories/textual/tree_text_memory/retrieve/xinyusearch.py +++ b/src/memos/memories/textual/tree_text_memory/retrieve/xinyusearch.py @@ -3,13 +3,15 @@ import json import uuid +from concurrent.futures import ThreadPoolExecutor, as_completed from datetime import datetime import requests from memos.embedders.factory import OllamaEmbedder from memos.log import get_logger -from memos.memories.textual.item import TextualMemoryItem, TreeNodeTextualMemoryMetadata +from memos.mem_reader.base import BaseMemReader +from memos.memories.textual.item import TextualMemoryItem logger = get_logger(__name__) @@ -93,8 +95,8 @@ def search(self, query: str, max_results: int | None = None) -> list[dict]: "online_search": { "max_entries": max_results, "cache_switch": False, - "baidu_field": {"switch": True, "mode": "relevance", "type": "page"}, - "bing_field": {"switch": False, "mode": "relevance", "type": "page_web"}, + "baidu_field": {"switch": False, "mode": "relevance", "type": "page"}, + "bing_field": {"switch": True, "mode": "relevance", "type": "page"}, "sogou_field": {"switch": False, "mode": "relevance", "type": "page"}, }, "request_id": "memos" + str(uuid.uuid4()), @@ -112,6 +114,7 @@ def __init__( access_key: str, search_engine_id: str, embedder: OllamaEmbedder, + reader: BaseMemReader, max_results: int = 20, ): """ @@ -121,12 +124,14 @@ def __init__( access_key: Xinyu API access key embedder: Embedder instance for generating embeddings max_results: Maximum number of results to retrieve + reader: MemReader Moduel to deal with internet contents """ self.xinyu_api = XinyuSearchAPI(access_key, search_engine_id, max_results=max_results) self.embedder = embedder + self.reader = reader def retrieve_from_internet( - self, query: str, top_k: int = 10, parsed_goal=None + self, query: str, top_k: int = 10, parsed_goal=None, info=None ) -> list[TextualMemoryItem]: """ Retrieve information from Xinyu search and convert to TextualMemoryItem format @@ -135,7 +140,7 @@ def retrieve_from_internet( query: Search query top_k: Number of results to return parsed_goal: Parsed task goal (optional) - + info (dict): Leave a record of memory consumption. Returns: List of TextualMemoryItem """ @@ -143,63 +148,25 @@ def retrieve_from_internet( search_results = self.xinyu_api.search(query, max_results=top_k) # Convert to TextualMemoryItem format - memory_items = [] - - for _, result in enumerate(search_results): - # Extract basic information from Xinyu response format - title = result.get("title", "") - content = result.get("content", "") - summary = result.get("summary", "") - url = result.get("url", "") - publish_time = result.get("publish_time", "") - if publish_time: + memory_items: list[TextualMemoryItem] = [] + + with ThreadPoolExecutor(max_workers=8) as executor: + futures = [ + executor.submit(self._process_result, result, query, parsed_goal, info) + for result in search_results + ] + for future in as_completed(futures): try: - publish_time = datetime.strptime(publish_time, "%Y-%m-%d %H:%M:%S").strftime( - "%Y-%m-%d" - ) + memory_items.extend(future.result()) except Exception as e: - logger.error(f"xinyu search error: {e}") - publish_time = datetime.now().strftime("%Y-%m-%d") - else: - publish_time = datetime.now().strftime("%Y-%m-%d") - source = result.get("source", "") - site = result.get("site", "") - if site: - site = site.split("|")[0] - - # Combine memory content - memory_content = ( - f"Title: {title}\nSummary: {summary}\nContent: {content[:200]}...\nSource: {url}" - ) + logger.error(f"Error processing search result: {e}") - # Create metadata - metadata = TreeNodeTextualMemoryMetadata( - user_id=None, - session_id=None, - status="activated", - type="fact", # Search results are usually factual information - memory_time=publish_time, - source="web", - confidence=85.0, # Confidence level for search information - entities=self._extract_entities(title, content, summary), - tags=self._extract_tags(title, content, summary, parsed_goal), - visibility="public", - memory_type="LongTermMemory", # Search results as working memory - key=title, - sources=[url] if url else [], - embedding=self.embedder.embed([memory_content])[0], - created_at=datetime.now().isoformat(), - usage=[], - background=f"Xinyu search result from {site or source}", - ) - # Create TextualMemoryItem - memory_item = TextualMemoryItem( - id=str(uuid.uuid4()), memory=memory_content, metadata=metadata - ) + unique_memory_items = {} + for item in memory_items: + if item.memory not in unique_memory_items: + unique_memory_items[item.memory] = item - memory_items.append(memory_item) - - return memory_items + return list(unique_memory_items.values()) def _extract_entities(self, title: str, content: str, summary: str) -> list[str]: """ @@ -333,3 +300,38 @@ def _extract_tags(self, title: str, content: str, summary: str, parsed_goal=None tags.extend(parsed_goal.tags) return list(set(tags))[:15] # Limit to 15 tags + + def _process_result( + self, result: dict, query: str, parsed_goal: str, info: dict + ) -> list[TextualMemoryItem]: + title = result.get("title", "") + content = result.get("content", "") + summary = result.get("summary", "") + url = result.get("url", "") + publish_time = result.get("publish_time", "") + if publish_time: + try: + publish_time = datetime.strptime(publish_time, "%Y-%m-%d %H:%M:%S").strftime( + "%Y-%m-%d" + ) + except Exception as e: + logger.error(f"xinyu search error: {e}") + publish_time = datetime.now().strftime("%Y-%m-%d") + else: + publish_time = datetime.now().strftime("%Y-%m-%d") + + read_items = self.reader.get_memory([content], type="doc", info=info) + + memory_items = [] + for read_item_i in read_items[0]: + read_item_i.memory = ( + f"Title: {title}\nNewsTime: {publish_time}\nSummary: {summary}\n" + f"Content: {read_item_i.memory}" + ) + read_item_i.metadata.source = "web" + read_item_i.metadata.memory_type = "OuterMemory" + read_item_i.metadata.sources = [url] if url else [] + read_item_i.metadata.visibility = "public" + + memory_items.append(read_item_i) + return memory_items diff --git a/src/memos/memos_tools/dinding_report_bot.py b/src/memos/memos_tools/dinding_report_bot.py new file mode 100644 index 000000000..9791cf65a --- /dev/null +++ b/src/memos/memos_tools/dinding_report_bot.py @@ -0,0 +1,422 @@ +"""dinding_report_bot.py""" + +import base64 +import contextlib +import hashlib +import hmac +import json +import os +import time +import urllib.parse + +from datetime import datetime +from uuid import uuid4 + +from dotenv import load_dotenv + + +load_dotenv() + +try: + import io + + import matplotlib + import matplotlib.font_manager as fm + import numpy as np + import oss2 + import requests + + from PIL import Image, ImageDraw, ImageFont + + matplotlib.use("Agg") + from alibabacloud_dingtalk.robot_1_0 import models as robot_models + from alibabacloud_dingtalk.robot_1_0.client import Client as DingtalkRobotClient + from alibabacloud_tea_openapi import models as open_api_models + from alibabacloud_tea_util import models as util_models +except ImportError as e: + raise ImportError( + f"DingDing bot dependencies not found: {e}. " + "Please install required packages: pip install requests oss2 pillow matplotlib alibabacloud-dingtalk" + ) from e + +# ========================= +# 🔧 common tools +# ========================= +ACCESS_TOKEN_USER = os.getenv("DINGDING_ACCESS_TOKEN_USER") +SECRET_USER = os.getenv("DINGDING_SECRET_USER") +ACCESS_TOKEN_ERROR = os.getenv("DINGDING_ACCESS_TOKEN_ERROR") +SECRET_ERROR = os.getenv("DINGDING_SECRET_ERROR") +OSS_CONFIG = { + "endpoint": os.getenv("OSS_ENDPOINT"), + "region": os.getenv("OSS_REGION"), + "bucket_name": os.getenv("OSS_BUCKET_NAME"), + "oss_access_key_id": os.getenv("OSS_ACCESS_KEY_ID"), + "oss_access_key_secret": os.getenv("OSS_ACCESS_KEY_SECRET"), + "public_base_url": os.getenv("OSS_PUBLIC_BASE_URL"), +} +ROBOT_CODE = os.getenv("DINGDING_ROBOT_CODE") +DING_APP_KEY = os.getenv("DINGDING_APP_KEY") +DING_APP_SECRET = os.getenv("DINGDING_APP_SECRET") + + +# Get access_token +def get_access_token(): + url = f"https://oapi.dingtalk.com/gettoken?appkey={DING_APP_KEY}&appsecret={DING_APP_SECRET}" + resp = requests.get(url) + return resp.json()["access_token"] + + +def _pick_font(size: int = 48) -> ImageFont.ImageFont: + """ + Try to find a font from the following candidates (macOS / Windows / Linux are common): + Helvetica → Arial → DejaVu Sans + If found, use truetype, otherwise return the default bitmap font. + """ + candidates = ["Helvetica", "Arial", "DejaVu Sans"] + for name in candidates: + try: + font_path = fm.findfont(name, fallback_to_default=False) + return ImageFont.truetype(font_path, size) + except Exception: + continue + # Cannot find truetype, fallback to default and manually scale up + bitmap = ImageFont.load_default() + return ImageFont.FreeTypeFont(bitmap.path, size) if hasattr(bitmap, "path") else bitmap + + +def make_header( + title: str, + subtitle: str, + size=(1080, 260), + colors=("#C8F6E1", "#E8F8F5"), # Stylish mint green → lighter green + fg="#00956D", +) -> bytes: + """ + Generate a "Notification" banner with green gradient and bold large text. + title: main title (suggested ≤ 35 characters) + subtitle: sub title (e.g. "Notification") + """ + + # Can be placed inside or outside make_header + def _text_wh(draw: ImageDraw.ImageDraw, text: str, font: ImageFont.ImageFont): + """ + return (width, height), compatible with both Pillow old version (textsize) and new version (textbbox) + """ + if hasattr(draw, "textbbox"): # Pillow ≥ 8.0 + left, top, right, bottom = draw.textbbox((0, 0), text, font=font) + return right - left, bottom - top + else: # Pillow < 10.0 + return draw.textsize(text, font=font) + + w, h = size + # --- 1) background gradient --- + g = np.linspace(0, 1, w) + grad = np.outer(np.ones(h), g) + rgb0 = tuple(int(colors[0].lstrip("#")[i : i + 2], 16) for i in (0, 2, 4)) + rgb1 = tuple(int(colors[1].lstrip("#")[i : i + 2], 16) for i in (0, 2, 4)) + img = np.zeros((h, w, 3), dtype=np.uint8) + for i in range(3): + img[:, :, i] = rgb0[i] * (1 - grad) + rgb1[i] * grad + im = Image.fromarray(img) + + # --- 2) text --- + draw = ImageDraw.Draw(im) + font_title = _pick_font(54) # main title + font_sub = _pick_font(30) # sub title + + # center alignment + title_w, title_h = _text_wh(draw, title, font_title) + sub_w, sub_h = _text_wh(draw, subtitle, font_sub) + + title_x = (w - title_w) // 2 + title_y = h // 2 - title_h + sub_x = (w - sub_w) // 2 + sub_y = title_y + title_h + 8 + + draw.text((title_x, title_y), title, fill=fg, font=font_title) + draw.text((sub_x, sub_y), subtitle, fill=fg, font=font_sub) + + # --- 3) PNG bytes --- + buf = io.BytesIO() + im.save(buf, "PNG") + return buf.getvalue() + + +def _sign(secret: str, ts: str): + s = f"{ts}\n{secret}" + return urllib.parse.quote_plus( + base64.b64encode(hmac.new(secret.encode(), s.encode(), hashlib.sha256).digest()) + ) + + +def _send_md(title: str, md: str, type="user", at=None): + if type == "user": + access_token = ACCESS_TOKEN_USER + secret = SECRET_USER + else: + access_token = ACCESS_TOKEN_ERROR + secret = SECRET_ERROR + ts = str(round(time.time() * 1000)) + url = ( + f"https://oapi.dingtalk.com/robot/send?access_token={access_token}" + f"×tamp={ts}&sign={_sign(secret, ts)}" + ) + payload = { + "msgtype": "markdown", + "markdown": {"title": title, "text": md}, + "at": at or {"atUserIds": [], "isAtAll": False}, + } + requests.post(url, headers={"Content-Type": "application/json"}, data=json.dumps(payload)) + + +# ------------------------- OSS ------------------------- +def upload_bytes_to_oss( + data: bytes, + oss_dir: str = "xcy-share/jfzt/", + filename: str | None = None, + keep_latest: int = 1, # Keep latest N files; 0 = delete all +) -> str: + """ + - If filename_prefix is provided, delete the older files in {oss_dir}/{prefix}_*.png, only keep the latest keep_latest files + - Always create __.png → ensure the URL is unique + """ + filename_prefix = filename + + conf = OSS_CONFIG + auth = oss2.Auth(conf["oss_access_key_id"], conf["oss_access_key_secret"]) + bucket = oss2.Bucket(auth, conf["endpoint"], conf["bucket_name"]) + + # ---------- delete old files ---------- + if filename_prefix and keep_latest >= 0: + prefix_path = f"{oss_dir.rstrip('/')}/{filename_prefix}_" + objs = bucket.list_objects(prefix=prefix_path).object_list + old_files = [(o.key, o.last_modified) for o in objs if o.key.endswith(".png")] + if old_files and len(old_files) > keep_latest: + # sort by last_modified from new to old + old_files.sort(key=lambda x: x[1], reverse=True) + to_del = [k for k, _ in old_files[keep_latest:]] + for k in to_del: + with contextlib.suppress(Exception): + bucket.delete_object(k) + + # ---------- upload new file ---------- + ts = int(time.time()) + uniq = uuid4().hex + prefix = f"{filename_prefix}_" if filename_prefix else "" + object_name = f"{oss_dir.rstrip('/')}/{prefix}{ts}_{uniq}.png" + bucket.put_object(object_name, data) + + return f"{conf['public_base_url'].rstrip('/')}/{object_name}" + + +# --------- Markdown Table Helper --------- +def _md_table(data: dict, is_error: bool = False) -> str: + """ + Render a dict to a DingTalk-compatible Markdown table + - Normal statistics: single row, multiple columns + - Error distribution: two columns, multiple rows (error information/occurrence count) + """ + if is_error: # {"error_info":{idx:val}, "occurrence_count":{idx:val}} + header = "| error | count |\n|---|---|" + rows = "\n".join( + f"| {err} | {cnt} |" + for err, cnt in zip(data["error"].values(), data["count"].values(), strict=False) + ) + return f"{header}\n{rows}" + + # normal statistics + header = "| " + " | ".join(data.keys()) + " |\n|" + "|".join(["---"] * len(data)) + "|" + row = "| " + " | ".join(map(str, data.values())) + " |" + return f"{header}\n{row}" + + +def upload_to_oss( + local_path: str, + oss_dir: str = "xcy-share/jfzt/", + filename: str | None = None, # ← Same addition +) -> str: + """Upload a local file to OSS, support overwrite""" + with open(local_path, "rb") as f: + return upload_bytes_to_oss(f.read(), oss_dir=oss_dir, filename=filename) + + +def send_ding_reminder( + access_token: str, robot_code: str, user_ids: list[str], content: str, remind_type: int = 0 +): + """ + :param access_token: DingTalk access_token (usually permanent when using a robot) + :param robot_code: Robot code applied on the open platform + :param user_ids: DingTalk user_id list + :param content: Message content to send + :param remind_type: 1=in-app notification, 2=phone reminder, 3=SMS reminder + """ + # initialize client + config = open_api_models.Config(protocol="https", region_id="central") + client = DingtalkRobotClient(config) + + # request headers + headers = robot_models.RobotSendDingHeaders(x_acs_dingtalk_access_token=access_token) + + # request body + req = robot_models.RobotSendDingRequest( + robot_code=robot_code, + remind_type=remind_type, + receiver_user_id_list=user_ids, + content=content, + ) + + # send + try: + client.robot_send_ding_with_options(req, headers, util_models.RuntimeOptions()) + print("✅ DING message sent successfully") + except Exception as e: + print("❌ DING message sent failed:", e) + + +def error_bot( + err: str, + title: str = "Error Alert", + level: str = "P2", # ← Add alert level + user_ids: list[str] | None = None, # ← @users in group +): + """ + send error alert + level can be set to P0 / P1 / P2, corresponding to red / orange / yellow + if title_color is provided, it will be overridden by level + """ + # ---------- Level → Color scheme & Emoji ---------- + level_map = { + "P0": {"color": "#C62828", "grad": ("#FFE4E4", "#FFD3D3"), "emoji": "🔴"}, + "P1": {"color": "#E65100", "grad": ("#FFE9D6", "#FFD7B5"), "emoji": "🟠"}, + "P2": {"color": "#EF6C00", "grad": ("#FFF6D8", "#FFECB5"), "emoji": "🟡"}, + } + lv = level.upper() + if lv not in level_map: + lv = "P0" # Default to P0 if invalid + style = level_map[lv] + + # If external title_color is specified, override with level color scheme + title_color = style["color"] + + # ---------- Generate gradient banner ---------- + banner_bytes = make_header( + title=f"Level {lv}", # Fixed English + subtitle="Error Alert", # Display level + colors=style["grad"], + fg=style["color"], + ) + banner_url = upload_bytes_to_oss( + banner_bytes, + filename=f"error_banner_{title}_{lv.lower()}.png", # Overwrite fixed file for each level + ) + + # ---------- Markdown ---------- + colored_title = f"{title}" + at_suffix = "" + if user_ids: + at_suffix = "\n\n" + " ".join([f"@{m}" for m in user_ids]) + + md = ( + f"![banner]({banner_url})\n\n" + f"### {style['emoji']} {colored_title}\n\n" + f"**Detail:**\n```\n{err}\n```\n" + # Visual indicator, pure color, no notification trigger + f"### 🔵 Attention:{at_suffix}\n\n" + f"Time: " + f"{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n" + ) + + # ---------- Send Markdown in group and @users ---------- + at_config = {"atUserIds": user_ids or [], "isAtAll": False} + _send_md(title, md, type="error", at=at_config) + + user_ids_for_ding = user_ids # DingTalk user_id list + message = f"{title}\nMemos system error, please handle immediately" + + token = get_access_token() + + send_ding_reminder( + access_token=token, + robot_code=ROBOT_CODE, + user_ids=user_ids_for_ding, + content=message, + remind_type=3 if level == "P0" else 1, # 1 in-app DING 2 SMS DING 3 phone DING + ) + + +# --------- online_bot --------- +# ---------- Convert dict → colored KV lines ---------- +def _kv_lines(d: dict, emoji: str = "", heading: str = "", heading_color: str = "#00956D") -> str: + """ + Returns: + ### 📅 Daily Summary + - **Request count:** 1364 + ... + """ + parts = [f"### {emoji} {heading}"] + parts += [f"- **{k}:** {v}" for k, v in d.items()] + return "\n".join(parts) + + +# -------------- online_bot(colored title version) ----------------- +def online_bot( + header_name: str, + sub_title_name: str, + title_color: str, + other_data1: dict, + other_data2: dict, + emoji: dict, +): + heading_color = "#00956D" # Green for subtitle + + # 0) Banner + banner_bytes = make_header(header_name, sub_title_name) + banner_url = upload_bytes_to_oss(banner_bytes, filename="online_report.png") + + # 1) Colored main title + colored_title = f"{header_name}" + + # 3) Markdown + md = "\n\n".join( + filter( + None, + [ + f"![banner]({banner_url})", + f"### 🙄 {colored_title}\n\n", + _kv_lines( + other_data1, + next(iter(emoji.keys())), + next(iter(emoji.values())), + heading_color=heading_color, + ), + _kv_lines( + other_data2, + list(emoji.keys())[1], + list(emoji.values())[1], + heading_color=heading_color, + ), + f"Time: " + f"{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n", + ], + ) + ) + + _send_md(colored_title, md, type="user") + + +if __name__ == "__main__": + other_data = { + "recent_overall_data": "what is memos", + "site_data": "**📊 Simulated content\nLa la la 320hahaha155", + } + + online_bot( + header_name="TextualMemory", # must in English + sub_title_name="Search", # must in English + title_color="#00956D", + other_data1={"Retrieval source 1": "This is plain text memory retrieval content blablabla"}, + other_data2=other_data, + emoji={"Plain text memory retrieval source": "😨", "Retrieval content": "🕰🐛"}, + ) + print("All messages sent successfully") diff --git a/src/memos/memos_tools/notification_service.py b/src/memos/memos_tools/notification_service.py new file mode 100644 index 000000000..e7db020f4 --- /dev/null +++ b/src/memos/memos_tools/notification_service.py @@ -0,0 +1,44 @@ +""" +Simple online_bot integration utility. +""" + +import logging + +from collections.abc import Callable + + +logger = logging.getLogger(__name__) + + +def get_online_bot_function() -> Callable | None: + """ + Get online_bot function if available, otherwise return None. + + Returns: + online_bot function if available, None otherwise + """ + try: + from memos.memos_tools.dinding_report_bot import online_bot + + logger.info("online_bot function loaded successfully") + return online_bot + except ImportError as e: + logger.warning(f"Failed to import online_bot: {e}, returning None") + return None + + +def get_error_bot_function() -> Callable | None: + """ + Get error_bot function if available, otherwise return None. + + Returns: + error_bot function if available, None otherwise + """ + try: + from memos.memos_tools.dinding_report_bot import error_bot + + logger.info("error_bot function loaded successfully") + return error_bot + except ImportError as e: + logger.warning(f"Failed to import error_bot: {e}, returning None") + return None diff --git a/src/memos/memos_tools/notification_utils.py b/src/memos/memos_tools/notification_utils.py new file mode 100644 index 000000000..390a9a556 --- /dev/null +++ b/src/memos/memos_tools/notification_utils.py @@ -0,0 +1,96 @@ +""" +Notification utilities for MemOS product. +""" + +import logging + +from collections.abc import Callable +from typing import Any + + +logger = logging.getLogger(__name__) + + +def send_online_bot_notification( + online_bot: Callable | None, + header_name: str, + sub_title_name: str, + title_color: str, + other_data1: dict[str, Any], + other_data2: dict[str, Any], + emoji: dict[str, str], +) -> None: + """ + Send notification via online_bot if available. + + Args: + online_bot: The online_bot function or None + header_name: Header name for the report + sub_title_name: Subtitle for the report + title_color: Title color + other_data1: First data dict + other_data2: Second data dict + emoji: Emoji configuration dict + """ + if online_bot is None: + return + + try: + online_bot( + header_name=header_name, + sub_title_name=sub_title_name, + title_color=title_color, + other_data1=other_data1, + other_data2=other_data2, + emoji=emoji, + ) + + logger.info(f"Online bot notification sent successfully: {header_name}") + + except Exception as e: + logger.warning(f"Failed to send online bot notification: {e}") + + +def send_error_bot_notification( + error_bot: Callable | None, + err: str, + title: str = "MemOS Error", + level: str = "P2", + user_ids: list | None = None, +) -> None: + """ + Send error alert if error_bot is available. + + Args: + error_bot: The error_bot function or None + err: Error message + title: Alert title + level: Alert level (P0, P1, P2) + user_ids: List of user IDs to notify + """ + if error_bot is None: + return + + try: + error_bot( + err=err, + title=title, + level=level, + user_ids=user_ids or [], + ) + logger.info(f"Error alert sent successfully: {title}") + except Exception as e: + logger.warning(f"Failed to send error alert: {e}") + + +# Keep backward compatibility +def send_error_alert( + error_bot: Callable | None, + error_message: str, + title: str = "MemOS Error", + level: str = "P2", +) -> None: + """ + Send error alert if error_bot is available (backward compatibility). + """ + send_error_bot_notification(error_bot, error_message, title, level) diff --git a/src/memos/settings.py b/src/memos/settings.py index d55a52e8a..3b3f05ebd 100644 --- a/src/memos/settings.py +++ b/src/memos/settings.py @@ -1,7 +1,9 @@ +import os + from pathlib import Path -MEMOS_DIR = Path.cwd() / ".memos" +MEMOS_DIR = Path(os.getenv("MEMOS_BASE_PATH", Path.cwd())) / ".memos" DEBUG = False # "memos" or "memos.submodules" ... to filter logs from specific packages diff --git a/tests/graph_dbs/test_nebular.py b/tests/graph_dbs/test_nebular.py deleted file mode 100644 index fa651ed13..000000000 --- a/tests/graph_dbs/test_nebular.py +++ /dev/null @@ -1,410 +0,0 @@ -import json -import os - -from datetime import datetime, timezone - -import numpy as np - -from dotenv import load_dotenv - -from memos.configs.embedder import EmbedderConfigFactory -from memos.configs.graph_db import GraphDBConfigFactory -from memos.embedders.factory import EmbedderFactory -from memos.graph_dbs.factory import GraphStoreFactory -from memos.memories.textual.item import TextualMemoryItem, TreeNodeTextualMemoryMetadata - - -load_dotenv() - -gpt_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"), - }, -} -nebular_config = { - "hosts": json.loads(os.getenv("NEBULAR_HOSTS", "localhost")), - "user_name": os.getenv("NEBULAR_USER", "root"), - "password": os.getenv("NEBULAR_PASSWORD", "xxxxxx"), - "space": "memory_graph", - "auto_create": True, - "embedding_dimension": 3072, - "use_multi_db": False, -} - - -embedder_config = EmbedderConfigFactory.model_validate(gpt_config) -embedder = EmbedderFactory.from_config(embedder_config) - - -def embed_memory_item(memory: str) -> list[float]: - embedding = embedder.embed([memory])[0] - embedding_np = np.array(embedding, dtype=np.float32) - embedding_list = embedding_np.tolist() - return embedding_list - - -now = datetime.now(timezone.utc).isoformat() -test_node1 = TextualMemoryItem( - memory="This is a test node", - 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("This is a test node"), - ), -) - -test_node2 = TextualMemoryItem( - memory="This is another test node", - 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("This is another test node"), - ), -) - - -def test_get_memory_count(): - config = GraphDBConfigFactory(backend="nebular", config=nebular_config) - graph = GraphStoreFactory.from_config(config) - graph.clear() - - mem = test_node1 - graph.add_node(mem.id, mem.memory, mem.metadata.model_dump(exclude_none=True)) - - count = graph.get_memory_count('"LongTermMemory"') # quoting string literal for Cypher - print("Memory Count:", count) - assert count == 1 - - -def test_count_nodes(): - graph = GraphStoreFactory.from_config( - GraphDBConfigFactory( - backend="nebular", - config=nebular_config, - ) - ) - graph.clear() - - # Insert two nodes - for i in range(2): - mem = TextualMemoryItem( - memory=f"Memory {i}", - 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(f"Memory {i}"), - ), - ) - graph.add_node(mem.id, mem.memory, mem.metadata.model_dump(exclude_none=True)) - - count = graph.count_nodes('"LongTermMemory"') - print("Node Count:", count) - assert count == 2 - - -def test_get_nodes(): - graph = GraphStoreFactory.from_config( - GraphDBConfigFactory(backend="nebular", config=nebular_config) - ) - graph.clear() - - mem = test_node1 - graph.add_node(mem.id, mem.memory, mem.metadata.model_dump(exclude_none=True)) - - nodes = graph.get_nodes([mem.id]) - assert len(nodes) == 1 - assert nodes[0]["properties"]["id"] == mem.id - - -def test_edge_exists(): - graph = GraphStoreFactory.from_config( - GraphDBConfigFactory(backend="nebular", config=nebular_config) - ) - graph.clear() - - topic = TextualMemoryItem( - memory="Edge topic", - 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("Edge topic"), - ), - ) - - concept = TextualMemoryItem( - memory="Edge concept", - 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("Edge concept"), - ), - ) - - graph.add_node(topic.id, topic.memory, topic.metadata.model_dump(exclude_none=True)) - graph.add_node(concept.id, concept.memory, concept.metadata.model_dump(exclude_none=True)) - graph.add_edge(topic.id, concept.id, type="RELATE_TO") - - assert graph.edge_exists(topic.id, concept.id, type="RELATE_TO", direction="OUTGOING") - - -def test_get_edges(): - graph = GraphStoreFactory.from_config( - GraphDBConfigFactory(backend="nebular", config=nebular_config) - ) - graph.clear() - - source = TextualMemoryItem( - memory="Source", - 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("Source"), - ), - ) - target = TextualMemoryItem( - memory="Target", - 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("Target"), - ), - ) - graph.add_node(source.id, source.memory, source.metadata.model_dump(exclude_none=True)) - graph.add_node(target.id, target.memory, target.metadata.model_dump(exclude_none=True)) - graph.add_edge(source.id, target.id, type="PARENT") - - edges = graph.get_edges(source.id, type="PARENT", direction="OUTGOING") - assert len(edges) == 1 - assert edges[0]["from"] == source.id - assert edges[0]["to"] == target.id - assert edges[0]["type"] == "PARENT" - - -def test_get_all_memory_items(): - graph = GraphStoreFactory.from_config( - GraphDBConfigFactory( - backend="nebular", - config=nebular_config, - ) - ) - graph.clear() - - # Insert 2 WorkingMemory items - for i in range(2): - mem = TextualMemoryItem( - memory=f"Memory {i}", - metadata=TreeNodeTextualMemoryMetadata( - memory_type="WorkingMemory", - key="Research Topic", - hierarchy_level="topic", - type="fact", - memory_time="2024-01-01", - status="activated", - visibility="public", - updated_at=now, - embedding=embed_memory_item(f"Memory {i}"), - ), - ) - graph.add_node(mem.id, mem.memory, mem.metadata.model_dump(exclude_none=True)) - - # Retrieve all memory items of type WorkingMemory - items = graph.get_all_memory_items("WorkingMemory") - assert len(items) == 2 - assert all(item["properties"]["memory_type"] == "WorkingMemory" for item in items) - - -def test_get_structure_optimization_candidates(): - graph = GraphStoreFactory.from_config( - GraphDBConfigFactory( - backend="nebular", - config=nebular_config, - ) - ) - graph.clear() - - # Insert one isolated node (no parent or child) - mem = TextualMemoryItem( - memory="Isolated memory", - 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("Isolated memory"), - ), - ) - graph.add_node(mem.id, mem.memory, mem.metadata.model_dump(exclude_none=True)) - - # Insert one node with empty background (and no edges) - mem2 = TextualMemoryItem( - memory="Empty background memory", - 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("Empty background memory"), - ), - ) - graph.add_node(mem2.id, mem2.memory, mem2.metadata.model_dump(exclude_none=True)) - - # Find optimization candidates - candidates = graph.get_structure_optimization_candidates("LongTermMemory") - print("Optimization candidates:", candidates) - assert any("Isolated memory" in c["memory"] for c in candidates) - assert any("Empty background memory" in c["memory"] for c in candidates) - - -def test_drop_database(): - config = GraphDBConfigFactory( - backend="nebular", - config=nebular_config, - ) - graph = GraphStoreFactory.from_config(config) - - # Create a dummy node - mem = TextualMemoryItem( - memory="Temp for drop DB", - 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("Temp for drop DB"), - ), - ) - graph.add_node(mem.id, mem.memory, mem.metadata.model_dump(exclude_none=True)) - - # Drop the database - graph.drop_database() - - # Attempting any operation afterward should raise an error or fail (optional) - try: - _ = graph.get_all_memory_items("WorkingMemory") - except Exception as e: - print("Expected exception after DB drop:", str(e)) - assert "Current working graph not found" in str(e) - - -def test_get_by_metadata(): - config = GraphDBConfigFactory( - backend="nebular", - config=nebular_config, - ) - graph = GraphStoreFactory.from_config(config) - graph.clear() - - mem1 = TextualMemoryItem( - memory="AI for science", - metadata=TreeNodeTextualMemoryMetadata( - memory_type="LongTermMemory", - key="AI Science", - confidence=92.5, - tags=["AI", "science"], - hierarchy_level="topic", - type="fact", - memory_time="2024-01-01", - status="activated", - visibility="public", - updated_at=now, - embedding=embed_memory_item("AI for science"), - ), - ) - mem2 = TextualMemoryItem( - memory="Neurosymbolic reasoning", - metadata=TreeNodeTextualMemoryMetadata( - memory_type="LongTermMemory", - key="Neurosymbolic", - tags=["symbolic", "reasoning"], - confidence=88.0, - hierarchy_level="topic", - type="fact", - memory_time="2024-01-01", - status="activated", - visibility="public", - updated_at=now, - embedding=embed_memory_item("Neurosymbolic reasoning"), - ), - ) - graph.add_node(mem1.id, mem1.memory, mem1.metadata.model_dump(exclude_none=True)) - graph.add_node(mem2.id, mem2.memory, mem2.metadata.model_dump(exclude_none=True)) - - # Exact match filter - result_ids = graph.get_by_metadata([{"field": "key", "op": "=", "value": '"AI Science"'}]) - assert mem1.id in result_ids - assert mem2.id not in result_ids - - # Confidence filter - result_ids = graph.get_by_metadata([{"field": "confidence", "op": ">=", "value": 90.0}]) - assert mem1.id in result_ids - assert mem2.id not in result_ids - - # Tag contains filter TODO - result_ids = graph.get_by_metadata([{"field": "tags", "op": "contains", "value": '["AI"]'}]) - assert mem1.id in result_ids - assert mem2.id not in result_ids - - # In set filter - result_ids = graph.get_by_metadata( - [{"field": "key", "op": "in", "value": '["AI Science", "Neurosymbolic"]'}] - ) - assert mem1.id in result_ids - assert mem2.id in result_ids diff --git a/tests/mem_os/test_memos_core.py b/tests/mem_os/test_memos_core.py index 88923fc94..0ebe49ff5 100644 --- a/tests/mem_os/test_memos_core.py +++ b/tests/mem_os/test_memos_core.py @@ -326,7 +326,16 @@ def test_search_memories( assert "para_mem" in result assert len(result["text_mem"]) == 1 assert result["text_mem"][0]["cube_id"] == "test_cube_1" - mock_mem_cube.text_mem.search.assert_called_once_with("football", top_k=5) + # Verify the search was called with the correct parameters + mock_mem_cube.text_mem.search.assert_called_once() + call_args = mock_mem_cube.text_mem.search.call_args + assert call_args[0] == ("football",) # positional args + assert call_args[1]["top_k"] == 5 + assert call_args[1]["mode"] == "fast" + assert call_args[1]["manual_close_internet"] + assert "info" in call_args[1] + assert call_args[1]["info"]["user_id"] == "test_user" + assert "session_id" in call_args[1]["info"] @patch("memos.mem_os.core.UserManager") @patch("memos.mem_os.core.MemReaderFactory") diff --git a/tests/mem_reader/test_simple_structure.py b/tests/mem_reader/test_simple_structure.py index 91996da91..18b674159 100644 --- a/tests/mem_reader/test_simple_structure.py +++ b/tests/mem_reader/test_simple_structure.py @@ -117,18 +117,18 @@ def test_get_scene_data_info_with_chat(self): self.assertEqual(len(result), 1) self.assertEqual(result[0][0], "user: [3 May 2025]: I'm feeling a bit down today.") - @patch("memos.parsers.factory.ParserFactory") + @patch("memos.mem_reader.simple_struct.ParserFactory") def test_get_scene_data_info_with_doc(self, mock_parser_factory): """Test parsing document files.""" parser_instance = MagicMock() parser_instance.parse.return_value = "Parsed document text.\n" mock_parser_factory.from_config.return_value = parser_instance - scene_data = ["tests/mem_reader/test.txt"] + scene_data = [{"fake_file_like": "should trigger parse"}] result = self.reader.get_scene_data_info(scene_data, type="doc") self.assertIsInstance(result, list) - self.assertEqual(result[0]["text"], "Parsed document text\n") + self.assertEqual(result[0]["text"], "Parsed document text.\n") def test_parse_json_result_success(self): """Test successful JSON parsing.""" diff --git a/tests/mem_scheduler/test_scheduler.py b/tests/mem_scheduler/test_scheduler.py index 30ccc934c..b35b7b174 100644 --- a/tests/mem_scheduler/test_scheduler.py +++ b/tests/mem_scheduler/test_scheduler.py @@ -49,8 +49,8 @@ def setUp(self): self.scheduler.mem_cube = self.mem_cube # Set current user and memory cube ID for testing - self.scheduler._current_user_id = "test_user" - self.scheduler._current_mem_cube_id = "test_cube" + self.scheduler.current_user_id = "test_user" + self.scheduler.current_mem_cube_id = "test_cube" def test_initialization(self): """Test that scheduler initializes with correct default values and handlers.""" diff --git a/tests/memories/textual/test_general.py b/tests/memories/textual/test_general.py index 8f5bf7966..94dcd5cd3 100644 --- a/tests/memories/textual/test_general.py +++ b/tests/memories/textual/test_general.py @@ -1,10 +1,8 @@ # TODO: Overcomplex. Use pytest fixtures instead of setUp/tearDown. -import json -import os import unittest import uuid -from unittest.mock import MagicMock, mock_open, patch +from unittest.mock import MagicMock, patch from memos.configs.embedder import EmbedderConfigFactory from memos.configs.llm import LLMConfigFactory diff --git a/tests/memories/textual/test_tree_searcher.py b/tests/memories/textual/test_tree_searcher.py index df7d1d773..729d7a4fc 100644 --- a/tests/memories/textual/test_tree_searcher.py +++ b/tests/memories/textual/test_tree_searcher.py @@ -94,7 +94,6 @@ def test_searcher_fine_mode_triggers_reasoner(mock_searcher): top_k=1, mode="fine", ) - assert mock_searcher.reasoner.reason.called assert len(result) == 1