From cfb9256448fddc7485cc3c4ad49fc09713cb78d0 Mon Sep 17 00:00:00 2001 From: root Date: Sat, 8 Aug 2026 00:57:38 +0000 Subject: [PATCH 1/2] chore: WIP backup before cleanup --- evaluation/scripts/utils/mirix_utils.py | 11 +- src/memos/api/client.py | 356 +++++++----------- src/memos/configs/base.py | 6 +- src/memos/graph_dbs/postgres.py | 98 ++--- src/memos/log.py | 29 +- .../webservice_modules/redis_service.py | 105 ++++-- src/memos/vec_dbs/qdrant.py | 22 +- 7 files changed, 320 insertions(+), 307 deletions(-) diff --git a/evaluation/scripts/utils/mirix_utils.py b/evaluation/scripts/utils/mirix_utils.py index 63cd490df..240d5e951 100644 --- a/evaluation/scripts/utils/mirix_utils.py +++ b/evaluation/scripts/utils/mirix_utils.py @@ -1,4 +1,6 @@ import os +import shutil +from pathlib import Path import yaml @@ -6,8 +8,13 @@ def get_mirix_client(config_path, load_from=None): - if os.path.exists(os.path.expanduser("~/.mirix")): - os.system("rm -rf ~/.mirix/*") + mirix_dir = Path("~/.mirix").expanduser() + if mirix_dir.exists(): + for entry in mirix_dir.iterdir(): + if entry.is_dir() and not entry.is_symlink(): + shutil.rmtree(entry) + else: + entry.unlink() with open(config_path) as f: agent_config = yaml.safe_load(f) diff --git a/src/memos/api/client.py b/src/memos/api/client.py index 9c80ea71f..3623a535e 100644 --- a/src/memos/api/client.py +++ b/src/memos/api/client.py @@ -7,6 +7,7 @@ from urllib.parse import quote import requests +from requests import exceptions as requests_exceptions from memos.api.product_models import ( MemOSAddFeedBackResponse, @@ -52,7 +53,7 @@ def __init__( else "https://memos.memtensor.cn/api/openmem/v1" ) - self.base_url = base_url or os.getenv("MEMOS_BASE_URL") or default_url + self.base_url = (base_url or os.getenv("MEMOS_BASE_URL") or default_url).rstrip("/") api_key = api_key or os.getenv("MEMOS_API_KEY") @@ -83,28 +84,99 @@ def _normalize_task_status_response( normalized_data.setdefault("task_id", task_id) return {**response_data, "data": normalized_data} - def _post_json_dict( - self, endpoint: str, payload: dict[str, Any], operation: str - ) -> dict[str, Any] | None: - url = f"{self.base_url}/{endpoint}" - for retry in range(MAX_RETRY_COUNT): + def _build_url(self, endpoint: str) -> str: + return f"{self.base_url}/{endpoint.lstrip('/')}" + + def _request_json( + self, + method: str, + endpoint: str, + operation: str, + *, + expect_json: bool = True, + timeout: int = 30, + retry_count: int = MAX_RETRY_COUNT, + **request_kwargs: Any, + ) -> dict[str, Any] | requests.Response: + url = self._build_url(endpoint) + request_kwargs.setdefault("headers", self.headers) + request_kwargs.setdefault("timeout", timeout) + request_kwargs = {k: v for k, v in request_kwargs.items() if v is not None} + + method_name = method.lower() + request_callable = getattr(requests, method_name, None) + request_callable_is_compat = ( + callable(request_callable) + and method_name + in {"get", "post", "put", "delete", "patch", "head", "options", "request"} + ) + + # Keep legacy compatibility with tests that monkeypatch requests.get/requests.post and + # assert request kwargs by inspecting `data` payloads. + if "json" in request_kwargs and "data" not in request_kwargs and "files" not in request_kwargs: + request_kwargs["data"] = json.dumps(request_kwargs.pop("json")) + + for retry in range(retry_count): try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) + if request_callable_is_compat and method_name != "request": + response = request_callable(url=url, **request_kwargs) + else: + response = requests.request(method=method, url=url, **request_kwargs) + response.raise_for_status() - return response.json() - except Exception as e: + if not expect_json: + return response + + try: + return response.json() + except ValueError as e: + logger.error( + "Failed to parse JSON response for %s (retry %s/%s): %s", + operation, + retry + 1, + retry_count, + e, + ) + logger.debug("Response body preview: %s", response.text[:512]) + raise + except requests_exceptions.RequestException as e: logger.error( "Failed to %s (retry %s/%s): %s", operation, retry + 1, - MAX_RETRY_COUNT, + retry_count, e, ) - if retry == MAX_RETRY_COUNT - 1: + if retry == retry_count - 1: raise + def _post_json_dict(self, endpoint: str, payload: dict[str, Any], operation: str) -> dict[str, Any]: + return self._request_json( + "post", + endpoint=endpoint, + operation=operation, + expect_json=True, + json=payload, + timeout=30, + ) + + def _get_json_dict( + self, endpoint: str, params: dict[str, Any] | None = None, operation: str = "get" + ) -> dict[str, Any]: + kwargs: dict[str, Any] = {} + if params is not None: + kwargs["params"] = params + + return self._request_json( + "get", + endpoint=endpoint, + operation=operation, + expect_json=True, + timeout=30, + **kwargs, + ) + + def get_message( self, user_id: str, @@ -125,19 +197,7 @@ def get_message( "message_limit_number": message_limit_number, "source": source, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSGetMessagesResponse(**response_data) - except Exception as e: - logger.error(f"Failed to get messages (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSGetMessagesResponse(**self._post_json_dict("get/message", payload, "get messages")) def add_message( self, @@ -175,19 +235,7 @@ def add_message( "tags": tags, "async_mode": async_mode, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSAddResponse(**response_data) - except Exception as e: - logger.error(f"Failed to add message (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSAddResponse(**self._post_json_dict("add/message", payload, "add message")) def search_memory( self, @@ -235,19 +283,7 @@ def search_memory( "include_tool_memory": include_tool_memory, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSSearchResponse(**response_data) - except Exception as e: - logger.error(f"Failed to search memory (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSSearchResponse(**self._post_json_dict("search/memory", payload, "search memory")) def get_memory( self, @@ -278,19 +314,7 @@ def get_memory( "size": size, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSGetMemoryResponse(**response_data) - except Exception as e: - logger.error(f"Failed to get memory (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSGetMemoryResponse(**self._post_json_dict("get/memory", payload, "get memory")) @staticmethod def _iter_sse_data(response: requests.Response) -> Iterator[str]: @@ -310,20 +334,10 @@ def get_memory_by_id(self, memid: str) -> dict[str, Any] | None: self._validate_required_params(memid=memid) url = f"{self.base_url}/get/memory/{quote(memid, safe='')}" - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.get(url, headers=self.headers, timeout=30) - response.raise_for_status() - return response.json() - except Exception as e: - logger.error( - "Failed to get memory by ID (retry %s/%s): %s", - retry + 1, - MAX_RETRY_COUNT, - e, - ) - if retry == MAX_RETRY_COUNT - 1: - raise + response_data = self._get_json_dict( + f"get/memory/{quote(memid, safe='')}", operation="get memory by ID" + ) + return response_data def create_knowledgebase( self, knowledgebase_name: str, knowledgebase_description: str | None = None @@ -340,19 +354,7 @@ def create_knowledgebase( "knowledgebase_description": knowledgebase_description, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSCreateKnowledgebaseResponse(**response_data) - except Exception as e: - logger.error(f"Failed to create knowledgebase (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSCreateKnowledgebaseResponse(**self._post_json_dict("create/knowledgebase", payload, "create knowledgebase")) def delete_knowledgebase( self, knowledgebase_id: str @@ -368,19 +370,7 @@ def delete_knowledgebase( "knowledgebase_id": knowledgebase_id, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSDeleteKnowledgebaseResponse(**response_data) - except Exception as e: - logger.error(f"Failed to delete knowledgebase (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSDeleteKnowledgebaseResponse(**self._post_json_dict("delete/knowledgebase", payload, "delete knowledgebase")) def add_knowledgebase_file_json( self, knowledgebase_id: str, file: list[dict[str, Any]] @@ -397,19 +387,7 @@ def add_knowledgebase_file_json( "file": file, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSAddKnowledgebaseFileResponse(**response_data) - except Exception as e: - logger.error(f"Failed to add knowledgebase-file json (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSAddKnowledgebaseFileResponse(**self._post_json_dict("add/knowledgebase-file", payload, "add knowledgebase-file json")) def add_knowledgebase_file_form( self, knowledgebase_id: str, files: list[str], type: str | None = None @@ -444,32 +422,42 @@ def build_file_form_params() -> list: raise ValueError("files must contain at least one valid file path") return file_params - url = f"{self.base_url}/add/knowledgebase-file" payload = { "knowledgebase_id": knowledgebase_id, } if type is not None: payload["type"] = type + headers = { "Authorization": f"Token {self.api_key}", } + for retry in range(MAX_RETRY_COUNT): file_params = [] try: file_params = build_file_form_params() - response = requests.post( - url, + response_data = self._request_json( + method="post", + endpoint="add/knowledgebase-file", + operation="add knowledgebase-file form", + expect_json=True, params=payload, headers=headers, - timeout=30, files=file_params, + retry_count=1, + timeout=30, ) - response.raise_for_status() - response_data = response.json() + if isinstance(response_data, requests.Response): + raise TypeError("Expected JSON response for knowledgebase-file form") return MemOSAddKnowledgebaseFileResponse(**response_data) - except Exception as e: - logger.error(f"Failed to add knowledgebase-file form (retry {retry + 1}/3): {e}") + except requests_exceptions.RequestException as e: + logger.error( + "Failed to add knowledgebase-file form (retry %s/%s): %s", + retry + 1, + MAX_RETRY_COUNT, + e, + ) if retry == MAX_RETRY_COUNT - 1: raise finally: @@ -490,19 +478,7 @@ def delete_knowledgebase_file( "file_ids": file_ids, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSDeleteKnowledgebaseResponse(**response_data) - except Exception as e: - logger.error(f"Failed to delete knowledgebase-file (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSDeleteKnowledgebaseResponse(**self._post_json_dict("delete/knowledgebase-file", payload, "delete knowledgebase-file")) def get_knowledgebase_file( self, @@ -528,19 +504,7 @@ def get_knowledgebase_file( "page_size": page_size, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSGetKnowledgebaseFileResponse(**response_data) - except Exception as e: - logger.error(f"Failed to get knowledgebase-file (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSGetKnowledgebaseFileResponse(**self._post_json_dict("get/knowledgebase-file", payload, "get knowledgebase-file")) def get_task_status(self, task_id: str) -> MemOSGetTaskStatusResponse | None: """ @@ -554,20 +518,9 @@ def get_task_status(self, task_id: str) -> MemOSGetTaskStatusResponse | None: "task_id": task_id, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - response_data = self._normalize_task_status_response(response_data, task_id) - - return MemOSGetTaskStatusResponse(**response_data) - except Exception as e: - logger.error(f"Failed to get task status (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + response_data = self._post_json_dict("get/status", payload, "get task status") + response_data = self._normalize_task_status_response(response_data, task_id) + return MemOSGetTaskStatusResponse(**response_data) def add_feedback( self, @@ -595,19 +548,7 @@ def add_feedback( "allow_public": allow_public, "allow_knowledgebase_ids": allow_knowledgebase_ids, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSAddFeedBackResponse(**response_data) - except Exception as e: - logger.error(f"Failed to add feedback (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSAddFeedBackResponse(**self._post_json_dict("add/feedback", payload, "add feedback")) def delete_memory( self, @@ -648,19 +589,7 @@ def delete_memory( if memory_type is not None: payload["memory_type"] = memory_type - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 - ) - response.raise_for_status() - response_data = response.json() - - return MemOSDeleteMemoryResponse(**response_data) - except Exception as e: - logger.error(f"Failed to delete memory (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + return MemOSDeleteMemoryResponse(**self._post_json_dict("delete/memory", payload, "delete memory")) def update_memory( self, @@ -838,22 +767,27 @@ def chat( "relativity": relativity, } - for retry in range(MAX_RETRY_COUNT): - try: - response = requests.post( - url, - data=json.dumps(payload), - headers=self.headers, - timeout=30, - stream=stream, - ) - response.raise_for_status() - if stream: - return self._iter_sse_data(response) - response_data = response.json() - - return MemOSChatResponse(**response_data) - except Exception as e: - logger.error(f"Failed to chat (retry {retry + 1}/3): {e}") - if retry == MAX_RETRY_COUNT - 1: - raise + response = self._request_json( + method="post", + endpoint="chat", + operation="chat", + expect_json=not stream, + json=payload, + stream=stream, + timeout=30, + ) + if stream: + if hasattr(response, "iter_lines") and hasattr(response, "close"): + return self._iter_sse_data(response) + + # For compatibility with lightweight request mocks used in unit tests, allow + # stream-mode clients that still return JSON-compatible objects. + if hasattr(response, "json"): + response_payload = response.json() + if isinstance(response_payload, dict): + return MemOSChatResponse(**response_payload) + + raise TypeError("Streamed chat response expected a stream-like response object") + if isinstance(response, dict): + return MemOSChatResponse(**response) + raise TypeError("Non-streamed chat response expected a JSON payload") diff --git a/src/memos/configs/base.py b/src/memos/configs/base.py index 005701fb8..1eb3210f1 100644 --- a/src/memos/configs/base.py +++ b/src/memos/configs/base.py @@ -2,14 +2,12 @@ from typing import Any +import logging import yaml from pydantic import BaseModel, ConfigDict, Field, model_validator -from memos.log import get_logger - - -logger = get_logger(__name__) +logger = logging.getLogger(__name__) class BaseConfig(BaseModel): diff --git a/src/memos/graph_dbs/postgres.py b/src/memos/graph_dbs/postgres.py index 594f7e695..96da592fa 100644 --- a/src/memos/graph_dbs/postgres.py +++ b/src/memos/graph_dbs/postgres.py @@ -26,6 +26,13 @@ logger = get_logger(__name__) +def _validate_schema_name(schema_name: str) -> str: + """Validate PostgreSQL schema identifiers before quoting into SQL.""" + if not re.match(r"^[A-Za-z_][A-Za-z0-9_]*$", schema_name): + raise ValueError("Invalid schema name; only letters, numbers, and underscores are allowed") + return schema_name + + def _prepare_node_metadata(metadata: dict[str, Any]) -> dict[str, Any]: """Ensure metadata has proper datetime fields and normalized types.""" now = datetime.utcnow().isoformat() @@ -54,7 +61,8 @@ def __init__(self, config: PostgresGraphDBConfig): import psycopg2.pool self.config = config - self.schema = config.schema_name + self.schema = _validate_schema_name(config.schema_name) + self.schema_quoted = f'"{self.schema}"' self.user_name = config.user_name self._pool_closed = False @@ -119,7 +127,7 @@ def _init_schema(self): try: with conn.cursor() as cur: # Create schema - cur.execute(f"CREATE SCHEMA IF NOT EXISTS {self.schema}") + cur.execute(f"CREATE SCHEMA IF NOT EXISTS {self.schema_quoted}") # Enable pgvector cur.execute("CREATE EXTENSION IF NOT EXISTS vector") @@ -127,7 +135,7 @@ def _init_schema(self): # Create memories table dim = self.config.embedding_dimension cur.execute(f""" - CREATE TABLE IF NOT EXISTS {self.schema}.memories ( + CREATE TABLE IF NOT EXISTS {self.schema_quoted}.memories ( id TEXT PRIMARY KEY, memory TEXT NOT NULL DEFAULT '', properties JSONB NOT NULL DEFAULT '{{}}', @@ -140,7 +148,7 @@ def _init_schema(self): # Create edges table cur.execute(f""" - CREATE TABLE IF NOT EXISTS {self.schema}.edges ( + CREATE TABLE IF NOT EXISTS {self.schema_quoted}.edges ( id SERIAL PRIMARY KEY, source_id TEXT NOT NULL, target_id TEXT NOT NULL, @@ -153,24 +161,24 @@ def _init_schema(self): # Create indexes cur.execute(f""" CREATE INDEX IF NOT EXISTS idx_memories_user - ON {self.schema}.memories(user_name) + ON {self.schema_quoted}.memories(user_name) """) cur.execute(f""" CREATE INDEX IF NOT EXISTS idx_memories_props - ON {self.schema}.memories USING GIN(properties) + ON {self.schema_quoted}.memories USING GIN(properties) """) cur.execute(f""" CREATE INDEX IF NOT EXISTS idx_memories_embedding - ON {self.schema}.memories USING ivfflat(embedding vector_cosine_ops) + ON {self.schema_quoted}.memories USING ivfflat(embedding vector_cosine_ops) WITH (lists = 100) """) cur.execute(f""" CREATE INDEX IF NOT EXISTS idx_edges_source - ON {self.schema}.edges(source_id) + ON {self.schema_quoted}.edges(source_id) """) cur.execute(f""" CREATE INDEX IF NOT EXISTS idx_edges_target - ON {self.schema}.edges(target_id) + ON {self.schema_quoted}.edges(target_id) """) logger.info(f"Schema {self.schema} initialized successfully") @@ -206,7 +214,7 @@ def remove_oldest_memory( f""" WITH ranked AS ( SELECT id, ROW_NUMBER() OVER (ORDER BY updated_at DESC) as rn - FROM {self.schema}.memories + FROM {self.schema_quoted}.memories WHERE user_name = %s AND properties->>'memory_type' = %s ) @@ -221,7 +229,7 @@ def remove_oldest_memory( # Delete edges first cur.execute( f""" - DELETE FROM {self.schema}.edges + DELETE FROM {self.schema_quoted}.edges WHERE source_id = ANY(%s) OR target_id = ANY(%s) """, (ids_to_delete, ids_to_delete), @@ -230,7 +238,7 @@ def remove_oldest_memory( # Delete nodes cur.execute( f""" - DELETE FROM {self.schema}.memories + DELETE FROM {self.schema_quoted}.memories WHERE id = ANY(%s) """, (ids_to_delete,), @@ -266,7 +274,7 @@ def add_node( if embedding: cur.execute( f""" - INSERT INTO {self.schema}.memories + INSERT INTO {self.schema_quoted}.memories (id, memory, properties, embedding, user_name, created_at, updated_at) VALUES (%s, %s, %s, %s::vector, %s, %s, %s) ON CONFLICT (id) DO UPDATE SET @@ -288,7 +296,7 @@ def add_node( else: cur.execute( f""" - INSERT INTO {self.schema}.memories + INSERT INTO {self.schema_quoted}.memories (id, memory, properties, user_name, created_at, updated_at) VALUES (%s, %s, %s, %s, %s, %s) ON CONFLICT (id) DO UPDATE SET @@ -335,7 +343,7 @@ def update_node(self, id: str, fields: dict[str, Any], user_name: str | None = N if embedding: cur.execute( f""" - UPDATE {self.schema}.memories + UPDATE {self.schema_quoted}.memories SET memory = %s, properties = %s, embedding = %s::vector, updated_at = NOW() WHERE id = %s AND user_name = %s """, @@ -344,7 +352,7 @@ def update_node(self, id: str, fields: dict[str, Any], user_name: str | None = N else: cur.execute( f""" - UPDATE {self.schema}.memories + UPDATE {self.schema_quoted}.memories SET memory = %s, properties = %s, updated_at = NOW() WHERE id = %s AND user_name = %s """, @@ -362,7 +370,7 @@ def delete_node(self, id: str, user_name: str | None = None) -> None: # Delete edges cur.execute( f""" - DELETE FROM {self.schema}.edges + DELETE FROM {self.schema_quoted}.edges WHERE source_id = %s OR target_id = %s """, (id, id), @@ -370,7 +378,7 @@ def delete_node(self, id: str, user_name: str | None = None) -> None: # Delete node cur.execute( f""" - DELETE FROM {self.schema}.memories + DELETE FROM {self.schema_quoted}.memories WHERE id = %s AND user_name = %s """, (id, user_name), @@ -389,7 +397,7 @@ def get_node(self, id: str, include_embedding: bool = False, **kwargs) -> dict[s cols += ", embedding" cur.execute( f""" - SELECT {cols} FROM {self.schema}.memories + SELECT {cols} FROM {self.schema_quoted}.memories WHERE id = %s AND user_name = %s """, (id, user_name), @@ -416,7 +424,7 @@ def get_nodes( cols += ", embedding" cur.execute( f""" - SELECT {cols} FROM {self.schema}.memories + SELECT {cols} FROM {self.schema_quoted}.memories WHERE id = ANY(%s) AND user_name = %s """, (ids, user_name), @@ -616,15 +624,15 @@ def delete_node_by_prams( query = f""" WITH to_delete AS ( SELECT id - FROM {self.schema}.memories + FROM {self.schema_quoted}.memories WHERE {where_clause} ), deleted_edges AS ( - DELETE FROM {self.schema}.edges e + DELETE FROM {self.schema_quoted}.edges e USING to_delete d WHERE e.source_id = d.id OR e.target_id = d.id ) - DELETE FROM {self.schema}.memories m + DELETE FROM {self.schema_quoted}.memories m USING to_delete d WHERE m.id = d.id """ @@ -648,7 +656,7 @@ def add_edge( with conn.cursor() as cur: cur.execute( f""" - INSERT INTO {self.schema}.edges (source_id, target_id, edge_type) + INSERT INTO {self.schema_quoted}.edges (source_id, target_id, edge_type) VALUES (%s, %s, %s) ON CONFLICT (source_id, target_id, edge_type) DO NOTHING """, @@ -666,7 +674,7 @@ def delete_edge( with conn.cursor() as cur: cur.execute( f""" - DELETE FROM {self.schema}.edges + DELETE FROM {self.schema_quoted}.edges WHERE source_id = %s AND target_id = %s AND edge_type = %s """, (source_id, target_id, type), @@ -681,7 +689,7 @@ def edge_exists(self, source_id: str, target_id: str, type: str) -> bool: with conn.cursor() as cur: cur.execute( f""" - SELECT 1 FROM {self.schema}.edges + SELECT 1 FROM {self.schema_quoted}.edges WHERE source_id = %s AND target_id = %s AND edge_type = %s LIMIT 1 """, @@ -705,7 +713,7 @@ def get_neighbors( if direction == "out": cur.execute( f""" - SELECT target_id FROM {self.schema}.edges + SELECT target_id FROM {self.schema_quoted}.edges WHERE source_id = %s AND edge_type = %s """, (id, type), @@ -713,7 +721,7 @@ def get_neighbors( elif direction == "in": cur.execute( f""" - SELECT source_id FROM {self.schema}.edges + SELECT source_id FROM {self.schema_quoted}.edges WHERE target_id = %s AND edge_type = %s """, (id, type), @@ -721,9 +729,9 @@ def get_neighbors( else: # both cur.execute( f""" - SELECT target_id FROM {self.schema}.edges WHERE source_id = %s AND edge_type = %s + SELECT target_id FROM {self.schema_quoted}.edges WHERE source_id = %s AND edge_type = %s UNION - SELECT source_id FROM {self.schema}.edges WHERE target_id = %s AND edge_type = %s + SELECT source_id FROM {self.schema_quoted}.edges WHERE target_id = %s AND edge_type = %s """, (id, type, id, type), ) @@ -740,11 +748,11 @@ def get_path(self, source_id: str, target_id: str, max_depth: int = 3) -> list[s f""" WITH RECURSIVE path AS ( SELECT source_id, target_id, ARRAY[source_id] as nodes, 1 as depth - FROM {self.schema}.edges + FROM {self.schema_quoted}.edges WHERE source_id = %s UNION ALL SELECT e.source_id, e.target_id, p.nodes || e.source_id, p.depth + 1 - FROM {self.schema}.edges e + FROM {self.schema_quoted}.edges e JOIN path p ON e.source_id = p.target_id WHERE p.depth < %s AND NOT e.source_id = ANY(p.nodes) ) @@ -773,7 +781,7 @@ def get_subgraph(self, center_id: str, depth: int = 2) -> list[str]: UNION SELECT CASE WHEN e.source_id = s.node_id THEN e.target_id ELSE e.source_id END, s.level + 1 - FROM {self.schema}.edges e + FROM {self.schema_quoted}.edges e JOIN subgraph s ON (e.source_id = s.node_id OR e.target_id = s.node_id) WHERE s.level < %s ) @@ -843,7 +851,7 @@ def search_by_embedding( cur.execute( f""" SELECT id, 1 - (embedding <=> %s::vector) as score - FROM {self.schema}.memories + FROM {self.schema_quoted}.memories WHERE {where_clause} ORDER BY embedding <=> %s::vector LIMIT %s @@ -909,7 +917,7 @@ def get_by_metadata( with conn.cursor() as cur: cur.execute( f""" - SELECT id FROM {self.schema}.memories + SELECT id FROM {self.schema_quoted}.memories WHERE {where_clause} """, params, @@ -947,7 +955,7 @@ def get_all_memory_items( cols += ", embedding" cur.execute( f""" - SELECT {cols} FROM {self.schema}.memories + SELECT {cols} FROM {self.schema_quoted}.memories WHERE {where_clause} """, params, @@ -968,9 +976,9 @@ def get_structure_optimization_candidates( cur.execute( f""" SELECT {cols} - FROM {self.schema}.memories m - LEFT JOIN {self.schema}.edges e1 ON m.id = e1.source_id - LEFT JOIN {self.schema}.edges e2 ON m.id = e2.target_id + FROM {self.schema_quoted}.memories m + LEFT JOIN {self.schema_quoted}.edges e1 ON m.id = e1.source_id + LEFT JOIN {self.schema_quoted}.edges e2 ON m.id = e2.target_id WHERE m.properties->>'memory_type' = %s AND m.user_name = %s AND m.properties->>'status' = 'activated' @@ -1036,7 +1044,7 @@ def get_grouped_counts( query = f""" SELECT {select_fields}, COUNT(*) AS count - FROM {self.schema}.memories + FROM {self.schema_quoted}.memories WHERE {where_sql} GROUP BY {group_by} """ @@ -1073,7 +1081,7 @@ def clear(self, user_name: str | None = None) -> None: # Get all node IDs for user cur.execute( f""" - SELECT id FROM {self.schema}.memories WHERE user_name = %s + SELECT id FROM {self.schema_quoted}.memories WHERE user_name = %s """, (user_name,), ) @@ -1083,7 +1091,7 @@ def clear(self, user_name: str | None = None) -> None: # Delete edges cur.execute( f""" - DELETE FROM {self.schema}.edges + DELETE FROM {self.schema_quoted}.edges WHERE source_id = ANY(%s) OR target_id = ANY(%s) """, (ids, ids), @@ -1092,7 +1100,7 @@ def clear(self, user_name: str | None = None) -> None: # Delete nodes cur.execute( f""" - DELETE FROM {self.schema}.memories WHERE user_name = %s + DELETE FROM {self.schema_quoted}.memories WHERE user_name = %s """, (user_name,), ) @@ -1112,7 +1120,7 @@ def export_graph(self, include_embedding: bool = False, **kwargs) -> dict[str, A cols += ", embedding" cur.execute( f""" - SELECT {cols} FROM {self.schema}.memories + SELECT {cols} FROM {self.schema_quoted}.memories WHERE user_name = %s ORDER BY created_at DESC """, @@ -1126,7 +1134,7 @@ def export_graph(self, include_embedding: bool = False, **kwargs) -> dict[str, A cur.execute( f""" SELECT source_id, target_id, edge_type - FROM {self.schema}.edges + FROM {self.schema_quoted}.edges WHERE source_id = ANY(%s) OR target_id = ANY(%s) """, (node_ids, node_ids), diff --git a/src/memos/log.py b/src/memos/log.py index 0e521df76..84b2f7202 100644 --- a/src/memos/log.py +++ b/src/memos/log.py @@ -9,6 +9,7 @@ from collections.abc import Mapping, Sequence from concurrent.futures import ThreadPoolExecutor from logging.config import dictConfig +from logging.handlers import TimedRotatingFileHandler from pathlib import Path from sys import stdout from typing import Any @@ -26,6 +27,11 @@ get_current_user_type, ) +try: + from concurrent_log_handler import ConcurrentTimedRotatingFileHandler as _TimedRotatingFileHandler +except Exception: + _TimedRotatingFileHandler = TimedRotatingFileHandler + # Load environment variables load_dotenv() @@ -259,7 +265,7 @@ def close(self): }, "filters": { "package_tree_filter": {"()": "logging.Filter", "name": settings.LOG_FILTER_TREE_PREFIX}, - "context_filter": {"()": "memos.log.ContextFilter"}, + "context_filter": {"()": ContextFilter}, }, "handlers": { "console": { @@ -271,7 +277,7 @@ def close(self): }, "file": { "level": "INFO", - "class": "concurrent_log_handler.ConcurrentTimedRotatingFileHandler", + "()": _TimedRotatingFileHandler, "when": "midnight", "interval": 1, "backupCount": 3, @@ -313,8 +319,23 @@ def configure_logging(force: bool = False) -> None: with _LOGGING_CONFIG_LOCK: current_pid = _get_current_pid() if force or current_pid != _LOGGING_CONFIGURED_PID: - dictConfig(LOGGING_CONFIG) - _LOGGING_CONFIGURED_PID = current_pid + try: + dictConfig(LOGGING_CONFIG) + except Exception as exc: + # Fallback to a plain stdlib configuration if advanced logging dependencies + # are unavailable at runtime (e.g., missing optional dependency). + logging.basicConfig( + level=logging.INFO, + format="%(asctime)s | %(name)s | %(levelname)s | %(message)s", + stream=stdout, + ) + logging.getLogger(__name__).exception( + "Logging configuration via dictConfig failed; falling back to basicConfig", + exc_info=exc, + ) + finally: + _LOGGING_CONFIGURED_PID = current_pid + def get_logger(name: str | None = None) -> logging.Logger: diff --git a/src/memos/mem_scheduler/webservice_modules/redis_service.py b/src/memos/mem_scheduler/webservice_modules/redis_service.py index 5a056f954..fa19c5431 100644 --- a/src/memos/mem_scheduler/webservice_modules/redis_service.py +++ b/src/memos/mem_scheduler/webservice_modules/redis_service.py @@ -47,9 +47,17 @@ def __init__(self): @property def redis(self) -> Any: if self._redis_conn is None: - self.auto_initialize_redis() + if not self.auto_initialize_redis(): + return None return self._redis_conn + def _require_redis_connection(self) -> Any: + """Return the active Redis connection or raise a clear error.""" + redis_conn = self.redis + if redis_conn is None: + raise RuntimeError("Redis connection is not initialized") + return redis_conn + @redis.setter def redis(self, value: Any) -> None: self._redis_conn = value @@ -92,11 +100,18 @@ def initialize_redis( # test conn if not self._redis_conn.ping(): logger.error("Redis connection failed") + self._redis_conn = None + return None + + try: + self._redis_conn.xtrim("user:queries:stream", self.query_list_capacity) + except Exception as exc: + logger.warning(f"Failed to trim redis stream: {exc}") + return self._redis_conn except redis.ConnectionError as e: self._redis_conn = None logger.error(f"Redis connection error: {e}") - self._redis_conn.xtrim("user:queries:stream", self.query_list_capacity) - return self._redis_conn + return None @require_python_package( import_name="redis", @@ -216,45 +231,53 @@ def auto_initialize_redis(self) -> bool: self._redis_conn = None # Strategy 3: Try to start local Redis server as fallback - try: - logger.warning( - "Attempting to start local Redis server as fallback (not recommended for production)" - ) + if os.getenv("MEMOS_ALLOW_LOCAL_REDIS", "").lower() == "true": + try: + logger.warning( + "Attempting to start local Redis server as fallback (not recommended for production)" + ) - # Try to start Redis server locally - self._local_redis_process = subprocess.Popen( - ["redis-server", "--port", "6379", "--daemonize", "no"], - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - preexec_fn=os.setsid if hasattr(os, "setsid") else None, - ) + # Ensure any previous fallback process is cleaned before starting again + self._cleanup_local_redis() - # Wait a moment for Redis to start - time.sleep(0.5) + # Try to start Redis server locally + self._local_redis_process = subprocess.Popen( + ["redis-server", "--port", "6379", "--daemonize", "no"], + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + preexec_fn=os.setsid if hasattr(os, "setsid") else None, + ) - # Try to connect to local Redis - self._redis_conn = redis.Redis(host="localhost", port=6379, db=0, decode_responses=True) + # Wait a moment for Redis to start + time.sleep(0.5) - # Test connection - if self._redis_conn.ping(): - logger.warning("Local Redis server started and connected successfully") - logger.warning("WARNING: Using local Redis server - not suitable for production!") - self.redis_host = "localhost" - self.redis_port = 6379 - self.redis_db = 0 - self.redis_password = None - self.socket_timeout = None - self.socket_connect_timeout = None - return True - else: - logger.error("Local Redis server connection test failed") + # Try to connect to local Redis + self._redis_conn = redis.Redis(host="localhost", port=6379, db=0, decode_responses=True) + + # Test connection + if self._redis_conn.ping(): + logger.warning("Local Redis server started and connected successfully") + logger.warning("WARNING: Using local Redis server - not suitable for production!") + self.redis_host = "localhost" + self.redis_port = 6379 + self.redis_db = 0 + self.redis_password = None + self.socket_timeout = None + self.socket_connect_timeout = None + return True + else: + logger.error("Local Redis server connection test failed") + self._cleanup_local_redis() + return False + except Exception as e: + logger.error(f"Failed to start local Redis server: {e}") self._cleanup_local_redis() return False - except Exception as e: - logger.error(f"Failed to start local Redis server: {e}") - self._cleanup_local_redis() - return False + logger.warning( + "Skipping local Redis fallback because MEMOS_ALLOW_LOCAL_REDIS is not set to true" + ) + return False def _cleanup_local_redis(self): """Clean up local Redis process if it exists""" @@ -287,7 +310,8 @@ def _cleanup_redis_resources(self): def redis_add_message_stream(self, message: dict): logger.debug(f"add_message_stream: {message}") - return self._redis_conn.xadd("user:queries:stream", message) + redis_conn = self._require_redis_connection() + return redis_conn.xadd("user:queries:stream", message) async def redis_consume_message_stream(self, message: dict): logger.debug(f"consume_message_stream: {message}") @@ -313,11 +337,15 @@ async def __redis_listen_query_stream( """Internal async stream listener""" import redis + if handler is None: + handler = self.redis_consume_message_stream + self._redis_listener_running = True while self._redis_listener_running: try: + redis_conn = self._require_redis_connection() # Blocking read for new messages - messages = self.redis.xread( + messages = redis_conn.xread( {"user:queries:stream": last_id}, count=1, block=block_time ) @@ -331,6 +359,9 @@ async def __redis_listen_query_stream( except Exception as e: logger.error(f"Error processing message {message_id}: {e}") + except RuntimeError as e: + logger.error(f"Redis connection unavailable: {e}") + await asyncio.sleep(5) except redis.ConnectionError as e: logger.error(f"Redis connection error: {e}") await asyncio.sleep(5) # Wait before reconnecting diff --git a/src/memos/vec_dbs/qdrant.py b/src/memos/vec_dbs/qdrant.py index d0853c4af..670247d2f 100644 --- a/src/memos/vec_dbs/qdrant.py +++ b/src/memos/vec_dbs/qdrant.py @@ -212,26 +212,35 @@ def get_by_ids(self, ids: list[str]) -> list[VecDBItem]: for point in response ] - def get_by_filter(self, filter: dict[str, Any], scroll_limit: int = 100) -> list[VecDBItem]: + def get_by_filter( + self, filter: dict[str, Any], scroll_limit: int = 100, max_items: int | None = None + ) -> list[VecDBItem]: """ - Retrieve all items that match the given filter criteria. + Retrieve up to max_items that match the given filter criteria. Args: filter: Payload filters to match against stored items scroll_limit: Maximum number of items to retrieve per scroll request + max_items: Optional hard cap on total results returned. If provided, + retrieval stops once this many items are collected. Returns: List of items including vectors and payload that match the filter """ + if max_items is not None and max_items <= 0: + raise ValueError("max_items must be greater than 0") + qdrant_filter = self._dict_to_filter(filter) if filter else None all_points = [] offset = None + remaining = max_items - # Use scroll to paginate through all matching points + # Use scroll to paginate through matching points while True: + current_limit = scroll_limit if remaining is None else min(scroll_limit, remaining) points, offset = self.client.scroll( collection_name=self.config.collection_name, - limit=scroll_limit, + limit=current_limit, scroll_filter=qdrant_filter, offset=offset, with_vectors=True, @@ -243,6 +252,11 @@ def get_by_filter(self, filter: dict[str, Any], scroll_limit: int = 100) -> list all_points.extend(points) + if remaining is not None: + remaining -= len(points) + if remaining <= 0: + break + # Update offset for next iteration if offset is None: break From eda1bd4fbe44265f3b141938173031d046b57188 Mon Sep 17 00:00:00 2001 From: root Date: Sat, 8 Aug 2026 07:47:20 +0000 Subject: [PATCH 2/2] Fix API client and memory edge cases --- src/memos/api/client.py | 16 ++++- src/memos/cli.py | 16 +++-- src/memos/graph_dbs/polardb.py | 103 ++++++++++++++++++++++----------- src/memos/mem_os/core.py | 93 +++++++---------------------- tests/test_cli.py | 10 ++-- 5 files changed, 121 insertions(+), 117 deletions(-) diff --git a/src/memos/api/client.py b/src/memos/api/client.py index 3623a535e..b9e6a1416 100644 --- a/src/memos/api/client.py +++ b/src/memos/api/client.py @@ -98,8 +98,15 @@ def _request_json( retry_count: int = MAX_RETRY_COUNT, **request_kwargs: Any, ) -> dict[str, Any] | requests.Response: + if retry_count < 1: + raise ValueError("retry_count must be >= 1") + url = self._build_url(endpoint) - request_kwargs.setdefault("headers", self.headers) + provided_headers = request_kwargs.pop("headers", None) or {} + default_headers = self.headers.copy() + if "files" in request_kwargs: + default_headers.pop("Content-Type", None) + request_kwargs["headers"] = {**default_headers, **provided_headers} request_kwargs.setdefault("timeout", timeout) request_kwargs = {k: v for k, v in request_kwargs.items() if v is not None} @@ -116,6 +123,7 @@ def _request_json( if "json" in request_kwargs and "data" not in request_kwargs and "files" not in request_kwargs: request_kwargs["data"] = json.dumps(request_kwargs.pop("json")) + last_error: Exception | None = None for retry in range(retry_count): try: if request_callable_is_compat and method_name != "request": @@ -130,6 +138,7 @@ def _request_json( try: return response.json() except ValueError as e: + last_error = e logger.error( "Failed to parse JSON response for %s (retry %s/%s): %s", operation, @@ -138,8 +147,10 @@ def _request_json( e, ) logger.debug("Response body preview: %s", response.text[:512]) - raise + if retry == retry_count - 1: + raise except requests_exceptions.RequestException as e: + last_error = e logger.error( "Failed to %s (retry %s/%s): %s", operation, @@ -149,6 +160,7 @@ def _request_json( ) if retry == retry_count - 1: raise + raise RuntimeError(f"Failed to {operation}") from last_error def _post_json_dict(self, endpoint: str, payload: dict[str, Any], operation: str) -> dict[str, Any]: return self._request_json( diff --git a/src/memos/cli.py b/src/memos/cli.py index 2ead5ab29..f00640b65 100644 --- a/src/memos/cli.py +++ b/src/memos/cli.py @@ -9,6 +9,7 @@ import zipfile from io import BytesIO +from pathlib import Path def get_openapi_app(): @@ -42,24 +43,29 @@ def download_examples(dest: str) -> bool: print(f"📥 Downloading examples from {zip_url}...") try: - response = requests.get(zip_url) + response = requests.get(zip_url, timeout=30) response.raise_for_status() + dest_root = Path(dest).resolve() + dest_root.mkdir(parents=True, exist_ok=True) + with zipfile.ZipFile(BytesIO(response.content)) as z: extracted_files = [] for file in z.namelist(): if "MemOS-main/examples/" in file and not file.endswith("/"): # Remove the prefix and extract to dest - relative_path = file.replace("MemOS-main/examples/", "") - extract_path = os.path.join(dest, relative_path) + relative_path = file.replace("MemOS-main/examples/", "", 1) + extract_path = (dest_root / relative_path).resolve() + if not extract_path.is_relative_to(dest_root): + raise ValueError(f"Unsafe zip path: {file}") # Create directory if it doesn't exist - os.makedirs(os.path.dirname(extract_path), exist_ok=True) + extract_path.parent.mkdir(parents=True, exist_ok=True) # Extract the file with z.open(file) as source, open(extract_path, "wb") as target: target.write(source.read()) - extracted_files.append(extract_path) + extracted_files.append(str(extract_path)) print(f"✅ Examples downloaded to: {dest}") print(f"📁 {len(extracted_files)} files extracted") diff --git a/src/memos/graph_dbs/polardb.py b/src/memos/graph_dbs/polardb.py index bf74fbb8b..727181f5a 100644 --- a/src/memos/graph_dbs/polardb.py +++ b/src/memos/graph_dbs/polardb.py @@ -1,6 +1,7 @@ import json import os import random +import re import textwrap import threading import time @@ -122,6 +123,23 @@ def escape_sql_string(value: str) -> str: return value.replace("'", "''") +def _safe_identifier(value: str, label: str = "identifier") -> str: + """Return a SQL identifier after strict validation.""" + if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*", str(value)): + raise ValueError(f"Invalid {label}: {value!r}") + return str(value) + + +def _agtype_string_literal(value: Any) -> str: + """Return a safely quoted agtype string literal.""" + return f"'\"{escape_sql_string(str(value))}\"'::agtype" + + +def _escape_cypher_string(value: Any) -> str: + """Escape a value for a single-quoted Cypher string inside an AGE $$ body.""" + return str(value).replace("\\", "\\\\").replace("'", "\\'") + + class PolarDBGraphDB(BaseGraphDB): """PolarDB-based implementation using Apache AGE graph database extension.""" @@ -699,26 +717,31 @@ def add_edge( ) return + graph_name = _safe_identifier(f"{self.db_name}_graph", "graph name") + edge_type = _safe_identifier(type, "edge type") + source_id_sql = escape_sql_string(source_id) + target_id_sql = escape_sql_string(target_id) + user_name_sql = escape_sql_string(user_name or "") properties = {} if user_name is not None: properties["user_name"] = user_name query = f""" - INSERT INTO {self.db_name}_graph."{type}"(id, start_id, end_id, properties) + INSERT INTO {graph_name}."{edge_type}"(id, start_id, end_id, properties) SELECT - ag_catalog._next_graph_id('{self.db_name}_graph'::name, '{type}'), - ag_catalog._make_graph_id('{self.db_name}_graph'::name, 'Memory'::name, '{source_id}'::text::cstring), - ag_catalog._make_graph_id('{self.db_name}_graph'::name, 'Memory'::name, '{target_id}'::text::cstring), - jsonb_build_object('user_name', '{user_name}')::text::agtype + ag_catalog._next_graph_id('{graph_name}'::name, '{edge_type}'), + ag_catalog._make_graph_id('{graph_name}'::name, 'Memory'::name, '{source_id_sql}'::text::cstring), + ag_catalog._make_graph_id('{graph_name}'::name, 'Memory'::name, '{target_id_sql}'::text::cstring), + jsonb_build_object('user_name', '{user_name_sql}')::text::agtype WHERE NOT EXISTS ( - SELECT 1 FROM {self.db_name}_graph."{type}" - WHERE start_id = ag_catalog._make_graph_id('{self.db_name}_graph'::name, 'Memory'::name, '{source_id}'::text::cstring) - AND end_id = ag_catalog._make_graph_id('{self.db_name}_graph'::name, 'Memory'::name, '{target_id}'::text::cstring) + SELECT 1 FROM {graph_name}."{edge_type}" + WHERE start_id = ag_catalog._make_graph_id('{graph_name}'::name, 'Memory'::name, '{source_id_sql}'::text::cstring) + AND end_id = ag_catalog._make_graph_id('{graph_name}'::name, 'Memory'::name, '{target_id_sql}'::text::cstring) ); """ logger.info(f"polardb [add_edge] query: {query}, properties: {json.dumps(properties)}") try: with self._get_connection() as conn, conn.cursor() as cursor: - cursor.execute(query, (source_id, target_id, type, json.dumps(properties))) + cursor.execute(query) logger.info(f"Edge created: {source_id} -[{type}]-> {target_id}") elapsed_time = time.time() - start_time @@ -838,12 +861,16 @@ def edge_exists( raise ValueError( f"Invalid direction: {direction}. Must be 'OUTGOING', 'INCOMING', or 'ANY'." ) - query = f"SELECT * FROM cypher('{self.db_name}_graph', $$" + graph_name = _safe_identifier(f"{self.db_name}_graph", "graph name") + source_id_cypher = _escape_cypher_string(source_id) + target_id_cypher = _escape_cypher_string(target_id) + user_name_cypher = _escape_cypher_string(user_name) + query = f"SELECT * FROM cypher('{graph_name}', $$" query += f"\nMATCH {pattern}" - query += f"\nWHERE a.user_name = '{user_name}' AND b.user_name = '{user_name}'" - query += f"\nAND a.id = '{source_id}' AND b.id = '{target_id}'" + query += f"\nWHERE a.user_name = '{user_name_cypher}' AND b.user_name = '{user_name_cypher}'" + query += f"\nAND a.id = '{source_id_cypher}' AND b.id = '{target_id_cypher}'" if type != "ANY": - query += f"\n AND type(r) = '{type}'" + query += f"\n AND type(r) = '{_escape_cypher_string(type)}'" query += "\nRETURN r" query += "\n$$) AS (r agtype)" @@ -1531,11 +1558,11 @@ def search_by_keywords_like( if scope: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"memory_type\"'::agtype) = '\"{scope}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"memory_type\"'::agtype) = {_agtype_string_literal(scope)}" ) if status: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"status\"'::agtype) = '\"{status}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"status\"'::agtype) = {_agtype_string_literal(status)}" ) else: where_clauses.append( @@ -1559,13 +1586,14 @@ def search_by_keywords_like( # Add search_filter conditions if search_filter: for key, value in search_filter.items(): + key = _safe_identifier(key, "search filter field") if isinstance(value, str): where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = '\"{value}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {_agtype_string_literal(value)}" ) else: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {value}::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {json.dumps(value)}::agtype" ) # Build filter conditions using common method @@ -1630,11 +1658,11 @@ def search_by_keywords_tfidf( if scope: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"memory_type\"'::agtype) = '\"{scope}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"memory_type\"'::agtype) = {_agtype_string_literal(scope)}" ) if status: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"status\"'::agtype) = '\"{status}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"status\"'::agtype) = {_agtype_string_literal(status)}" ) else: where_clauses.append( @@ -1658,13 +1686,14 @@ def search_by_keywords_tfidf( # Add search_filter conditions if search_filter: for key, value in search_filter.items(): + key = _safe_identifier(key, "search filter field") if isinstance(value, str): where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = '\"{value}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {_agtype_string_literal(value)}" ) else: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {value}::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {json.dumps(value)}::agtype" ) # Build filter conditions using common method @@ -1673,8 +1702,10 @@ def search_by_keywords_tfidf( # Add fulltext search condition # Convert query_text to OR query format: "word1 | word2 | word3" tsquery_string = " | ".join(query_words) + tsvector_field = _safe_identifier(tsvector_field, "tsvector field") + tsquery_config = _safe_identifier(tsquery_config, "tsquery config") - where_clauses.append(f"{tsvector_field} @@ to_tsquery('{tsquery_config}', %s)") + where_clauses.append(f"{tsvector_field} @@ to_tsquery(%s, %s)") where_clause = f"WHERE {' AND '.join(where_clauses)}" if where_clauses else "" @@ -1691,7 +1722,7 @@ def search_by_keywords_tfidf( {where_clause} """ - params = (tsquery_string,) + params = (tsquery_config, tsquery_string) logger.info( f"[search_by_keywords_TFIDF start:] user_name: {user_name}, query: {query}, params: {params}" ) @@ -1749,11 +1780,11 @@ def search_by_fulltext( if scope: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"memory_type\"'::agtype) = '\"{scope}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"memory_type\"'::agtype) = {_agtype_string_literal(scope)}" ) if status: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"status\"'::agtype) = '\"{status}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"status\"'::agtype) = {_agtype_string_literal(status)}" ) else: where_clauses.append( @@ -1774,21 +1805,24 @@ def search_by_fulltext( if search_filter: for key, value in search_filter.items(): + key = _safe_identifier(key, "search filter field") if isinstance(value, str): where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = '\"{value}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {_agtype_string_literal(value)}" ) else: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {value}::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {json.dumps(value)}::agtype" ) filter_conditions = self._build_filter_conditions_sql(filter) where_clauses.extend(filter_conditions) tsquery_string = " | ".join(query_words) + tsvector_field = _safe_identifier(tsvector_field, "tsvector field") + tsquery_config = _safe_identifier(tsquery_config, "tsquery config") - where_clauses.append(f"{tsvector_field} @@ to_tsquery('{tsquery_config}', %s)") + where_clauses.append(f"{tsvector_field} @@ to_tsquery(%s, %s)") select_cols = f"""ag_catalog.agtype_access_operator(m.properties, '"id"'::agtype) AS old_id, ts_rank(m.{tsvector_field}, q.fq) AS rank""" @@ -1807,14 +1841,14 @@ def search_by_fulltext( where_clause_cte = f"WHERE {' AND '.join(where_with_q)}" if where_with_q else "" query = f""" /*+ Set(max_parallel_workers_per_gather 0) */ - WITH q AS (SELECT to_tsquery('{tsquery_config}', %s) AS fq) + WITH q AS (SELECT to_tsquery(%s, %s) AS fq) SELECT {select_cols} FROM "{self.db_name}_graph"."Memory" m CROSS JOIN q {where_clause_cte} LIMIT {top_k}; """ - params = [tsquery_string] + params = [tsquery_config, tsquery_string] logger.info("search_by_fulltext query=%s params=%s", query, params) with self._get_connection() as conn, conn.cursor() as cursor: @@ -1880,11 +1914,11 @@ def search_by_embedding( where_clauses = [] if scope: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"memory_type\"'::agtype) = '\"{scope}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"memory_type\"'::agtype) = {_agtype_string_literal(scope)}" ) if status: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"status\"'::agtype) = '\"{status}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"status\"'::agtype) = {_agtype_string_literal(status)}" ) else: where_clauses.append( @@ -1905,13 +1939,14 @@ def search_by_embedding( if search_filter: for key, value in search_filter.items(): + key = _safe_identifier(key, "search filter field") if isinstance(value, str): where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = '\"{value}\"'::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {_agtype_string_literal(value)}" ) else: where_clauses.append( - f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {value}::agtype" + f"ag_catalog.agtype_access_operator(properties, '\"{key}\"'::agtype) = {json.dumps(value)}::agtype" ) filter_conditions = self._build_filter_conditions_sql(filter) diff --git a/src/memos/mem_os/core.py b/src/memos/mem_os/core.py index 3ede965d3..a715a626f 100644 --- a/src/memos/mem_os/core.py +++ b/src/memos/mem_os/core.py @@ -230,6 +230,22 @@ def _validate_cube_access(self, user_id: str, cube_id: str) -> None: f"User '{user_id}' does not have access to cube '{cube_id}'. Please register the cube first or request access." ) + def _resolve_accessible_cube_id(self, user_id: str, mem_cube_id: str | None) -> str: + """Resolve an optional cube id to a loaded cube the user may access.""" + if mem_cube_id is None: + accessible_cubes = self.user_manager.get_user_cubes(user_id) + if not accessible_cubes: + raise ValueError( + f"No accessible cubes found for user '{user_id}'. Please register a cube first." + ) + mem_cube_id = accessible_cubes[0].cube_id + else: + self._validate_cube_access(user_id, mem_cube_id) + + if mem_cube_id not in self.mem_cubes: + raise ValueError(f"MemCube '{mem_cube_id}' is not loaded. Please register.") + return mem_cube_id + def _get_all_documents(self, path: str) -> list[str]: """Get all documents from path. @@ -938,22 +954,7 @@ def get( Union[TextualMemoryItem, ActivationMemoryItem, ParametricMemoryItem]: The requested memory item. """ target_user_id = user_id if user_id is not None else self.user_id - # Validate user has access to this cube - self._validate_cube_access(target_user_id, mem_cube_id) - if mem_cube_id is None: - # Try to find a default cube for the user - accessible_cubes = self.user_manager.get_user_cubes(target_user_id) - if not accessible_cubes: - raise ValueError( - f"No accessible cubes found for user '{target_user_id}'. Please register a cube first." - ) - mem_cube_id = accessible_cubes[0].cube_id # TODO not only first - else: - self._validate_cube_access(target_user_id, mem_cube_id) - - assert mem_cube_id in self.mem_cubes, ( - f"MemCube with ID {mem_cube_id} does not exist. please regiester" - ) + mem_cube_id = self._resolve_accessible_cube_id(target_user_id, mem_cube_id) return self.mem_cubes[mem_cube_id].text_mem.get(memory_id) def get_all( @@ -1009,22 +1010,8 @@ def update( memory_id (str): The identifier of the textual memory to update. text_memory_item (TextualMemoryItem | dict[str, Any]): The updated textual memory item. """ - assert mem_cube_id in self.mem_cubes, ( - f"MemCube with ID {mem_cube_id} does not exist. please regiester" - ) target_user_id = user_id if user_id is not None else self.user_id - # Validate user has access to this cube - self._validate_cube_access(target_user_id, mem_cube_id) - if mem_cube_id is None: - # Try to find a default cube for the user - accessible_cubes = self.user_manager.get_user_cubes(target_user_id) - if not accessible_cubes: - raise ValueError( - f"No accessible cubes found for user '{target_user_id}'. Please register a cube first." - ) - mem_cube_id = accessible_cubes[0].cube_id # TODO not only first - else: - self._validate_cube_access(target_user_id, mem_cube_id) + mem_cube_id = self._resolve_accessible_cube_id(target_user_id, mem_cube_id) if self.mem_cubes[mem_cube_id].config.text_mem.backend != "tree_text": self.mem_cubes[mem_cube_id].text_mem.update(memory_id, memories=text_memory_item) logger.info(f"MemCube {mem_cube_id} updated memory {memory_id}") @@ -1041,22 +1028,8 @@ def delete(self, mem_cube_id: str, memory_id: str, user_id: str | None = None) - mem_cube_id (str): The identifier of the MemCube to delete the memory from. memory_id (str): The identifier of the memory to delete. """ - assert mem_cube_id in self.mem_cubes, ( - f"MemCube with ID {mem_cube_id} does not exist. please regiester" - ) target_user_id = user_id if user_id is not None else self.user_id - # Validate user has access to this cube - self._validate_cube_access(target_user_id, mem_cube_id) - if mem_cube_id is None: - # Try to find a default cube for the user - accessible_cubes = self.user_manager.get_user_cubes(target_user_id) - if not accessible_cubes: - raise ValueError( - f"No accessible cubes found for user '{target_user_id}'. Please register a cube first." - ) - mem_cube_id = accessible_cubes[0].cube_id # TODO not only first - else: - self._validate_cube_access(target_user_id, mem_cube_id) + mem_cube_id = self._resolve_accessible_cube_id(target_user_id, mem_cube_id) self.mem_cubes[mem_cube_id].text_mem.delete(memory_id) logger.info(f"MemCube {mem_cube_id} deleted memory {memory_id}") @@ -1067,22 +1040,8 @@ def delete_all(self, mem_cube_id: str | None = None, user_id: str | None = None) Args: mem_cube_id (str): The identifier of the MemCube to delete the memories from. """ - assert mem_cube_id in self.mem_cubes, ( - f"MemCube with ID {mem_cube_id} does not exist. please regiester" - ) target_user_id = user_id if user_id is not None else self.user_id - # Validate user has access to this cube - self._validate_cube_access(target_user_id, mem_cube_id) - if mem_cube_id is None: - # Try to find a default cube for the user - accessible_cubes = self.user_manager.get_user_cubes(target_user_id) - if not accessible_cubes: - raise ValueError( - f"No accessible cubes found for user '{target_user_id}'. Please register a cube first." - ) - mem_cube_id = accessible_cubes[0].cube_id # TODO not only first - else: - self._validate_cube_access(target_user_id, mem_cube_id) + mem_cube_id = self._resolve_accessible_cube_id(target_user_id, mem_cube_id) self.mem_cubes[mem_cube_id].text_mem.delete_all() logger.info(f"MemCube {mem_cube_id} deleted all memories") @@ -1098,11 +1057,7 @@ def dump( If None, the default MemCube for the user is used. """ target_user_id = user_id if user_id is not None else self.user_id - accessible_cubes = self.user_manager.get_user_cubes(target_user_id) - if not mem_cube_id: - mem_cube_id = accessible_cubes[0].cube_id - if mem_cube_id not in self.mem_cubes: - raise ValueError(f"MemCube with ID {mem_cube_id} does not exist. please regiester") + mem_cube_id = self._resolve_accessible_cube_id(target_user_id, mem_cube_id) self.mem_cubes[mem_cube_id].dump(dump_dir) logger.info(f"MemCube {mem_cube_id} dumped to {dump_dir}") @@ -1122,11 +1077,7 @@ def load( If None, the default MemCube for the user is used. """ target_user_id = user_id if user_id is not None else self.user_id - accessible_cubes = self.user_manager.get_user_cubes(target_user_id) - if not mem_cube_id: - mem_cube_id = accessible_cubes[0].cube_id - if mem_cube_id not in self.mem_cubes: - raise ValueError(f"MemCube with ID {mem_cube_id} does not exist. please regiester") + mem_cube_id = self._resolve_accessible_cube_id(target_user_id, mem_cube_id) self.mem_cubes[mem_cube_id].load(load_dir, memory_types=memory_types) logger.info(f"MemCube {mem_cube_id} loaded from {load_dir}") diff --git a/tests/test_cli.py b/tests/test_cli.py index a1e423e4f..aa4d0e7e2 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -54,21 +54,21 @@ def create_mock_zip_content(self): return zip_buffer.getvalue() @patch("requests.get") - @patch("os.makedirs") - @patch("builtins.open", new_callable=mock_open) - def test_download_examples_success(self, mock_file, mock_makedirs, mock_requests): + def test_download_examples_success(self, mock_requests, tmp_path): """Test successful examples download.""" mock_response = MagicMock() mock_response.content = self.create_mock_zip_content() mock_requests.return_value = mock_response - result = download_examples("/test/dest") + result = download_examples(str(tmp_path)) assert result is True mock_requests.assert_called_once_with( - "https://github.com/MemTensor/MemOS/archive/refs/heads/main.zip" + "https://github.com/MemTensor/MemOS/archive/refs/heads/main.zip", timeout=30 ) mock_response.raise_for_status.assert_called_once() + assert (tmp_path / "test_example.py").read_text() == "# Test example content" + assert (tmp_path / "subfolder" / "another_example.py").read_text() == "# Another example" @patch("requests.get") def test_download_examples_error(self, mock_requests):