From 05cd088b6e65c51f0149c6b660fcf3f4891e40be Mon Sep 17 00:00:00 2001 From: Kwizii Date: Sat, 22 Aug 2026 23:08:03 +0800 Subject: [PATCH 1/4] fix: normalize Postgres embedding strings for reorganizer --- src/memos/graph_dbs/postgres.py | 34 +++++++++-- .../test_postgres_embedding_parse.py | 56 +++++++++++++++++++ 2 files changed, 86 insertions(+), 4 deletions(-) create mode 100644 tests/graph_dbs/test_postgres_embedding_parse.py diff --git a/src/memos/graph_dbs/postgres.py b/src/memos/graph_dbs/postgres.py index 594f7e695..d7c678c42 100644 --- a/src/memos/graph_dbs/postgres.py +++ b/src/memos/graph_dbs/postgres.py @@ -26,16 +26,37 @@ logger = get_logger(__name__) +def _normalize_embedding_value(embedding: Any) -> list[float] | None: + """Coerce pgvector/psycopg2 embedding values into list[float] for GraphDBNode.""" + if embedding is None: + return None + if isinstance(embedding, list): + return [float(x) for x in embedding] + if isinstance(embedding, tuple): + return [float(x) for x in embedding] + if isinstance(embedding, str): + stripped = embedding.strip() + if not stripped: + return None + try: + parsed = json.loads(stripped) + except json.JSONDecodeError: + return None + if isinstance(parsed, list): + return [float(x) for x in parsed] + return None + return None + + def _prepare_node_metadata(metadata: dict[str, Any]) -> dict[str, Any]: """Ensure metadata has proper datetime fields and normalized types.""" now = datetime.utcnow().isoformat() metadata.setdefault("created_at", now) metadata.setdefault("updated_at", now) - # Normalize embedding type - embedding = metadata.get("embedding") - if embedding and isinstance(embedding, list): - metadata["embedding"] = [float(x) for x in embedding] + normalized = _normalize_embedding_value(metadata.get("embedding")) + if normalized is not None: + metadata["embedding"] = normalized return metadata @@ -437,6 +458,11 @@ def _parse_row(self, row, include_embedding: bool = False) -> dict[str, Any]: } if include_embedding and len(row) > 5: result["metadata"]["embedding"] = row[5] + normalized = _normalize_embedding_value(result["metadata"].get("embedding")) + if normalized is not None: + result["metadata"]["embedding"] = normalized + elif "embedding" in result["metadata"]: + del result["metadata"]["embedding"] return result @staticmethod diff --git a/tests/graph_dbs/test_postgres_embedding_parse.py b/tests/graph_dbs/test_postgres_embedding_parse.py new file mode 100644 index 000000000..b66f8404e --- /dev/null +++ b/tests/graph_dbs/test_postgres_embedding_parse.py @@ -0,0 +1,56 @@ +"""Regression tests for PostgresGraphDB embedding normalization.""" + +from __future__ import annotations + +import json + +from datetime import datetime +from typing import Any +from unittest.mock import MagicMock, patch + +from memos.graph_dbs.postgres import PostgresGraphDB, _normalize_embedding_value, _prepare_node_metadata + + +def test_normalize_embedding_value_parses_json_string() -> None: + assert _normalize_embedding_value("[0.5, -1.0]") == [0.5, -1.0] + + +def test_normalize_embedding_value_coerces_numeric_list() -> None: + assert _normalize_embedding_value([2, 3]) == [2.0, 3.0] + + +def test_normalize_embedding_value_rejects_invalid_string() -> None: + assert _normalize_embedding_value("not-json") is None + + +def test_prepare_node_metadata_normalizes_string_embedding() -> None: + metadata = _prepare_node_metadata({"embedding": "[0.25, 0.75]"}) + assert metadata["embedding"] == [0.25, 0.75] + + +def _build_db() -> PostgresGraphDB: + with ( + patch("memos.graph_dbs.postgres.require_python_package", lambda **kwargs: lambda fn: fn), + patch("psycopg2.pool.ThreadedConnectionPool", MagicMock()), + patch.object(PostgresGraphDB, "_init_schema", lambda self: None), + ): + config = MagicMock() + config.schema_name = "test_schema" + config.user_name = "user-1" + return PostgresGraphDB(config) + + +def test_parse_row_normalizes_string_embedding_from_vector_column() -> None: + db = _build_db() + row: tuple[Any, ...] = ( + "node-1", + "memory text", + json.dumps({"memory_type": "UserMemory"}), + datetime(2026, 8, 22, 12, 0, 0), + datetime(2026, 8, 22, 12, 0, 0), + "[0.25, 0.75]", + ) + + parsed = db._parse_row(row, include_embedding=True) + + assert parsed["metadata"]["embedding"] == [0.25, 0.75] From 0dc5db9a0fece9cb781f74f631d428ca879d0bea Mon Sep 17 00:00:00 2001 From: Kwizii Date: Sat, 22 Aug 2026 23:13:51 +0800 Subject: [PATCH 2/4] fix: harden Postgres embedding normalization per review --- src/memos/graph_dbs/postgres.py | 40 ++++++++++--------- .../test_postgres_embedding_parse.py | 20 ++++++++++ 2 files changed, 42 insertions(+), 18 deletions(-) diff --git a/src/memos/graph_dbs/postgres.py b/src/memos/graph_dbs/postgres.py index d7c678c42..8a4f07834 100644 --- a/src/memos/graph_dbs/postgres.py +++ b/src/memos/graph_dbs/postgres.py @@ -30,20 +30,23 @@ def _normalize_embedding_value(embedding: Any) -> list[float] | None: """Coerce pgvector/psycopg2 embedding values into list[float] for GraphDBNode.""" if embedding is None: return None - if isinstance(embedding, list): - return [float(x) for x in embedding] - if isinstance(embedding, tuple): - return [float(x) for x in embedding] - if isinstance(embedding, str): - stripped = embedding.strip() - if not stripped: - return None - try: - parsed = json.loads(stripped) - except json.JSONDecodeError: + try: + if isinstance(embedding, list): + return [float(x) for x in embedding] + if isinstance(embedding, tuple): + return [float(x) for x in embedding] + if isinstance(embedding, str): + stripped = embedding.strip() + if not stripped: + return None + try: + parsed = json.loads(stripped) + except json.JSONDecodeError: + return None + if isinstance(parsed, list): + return [float(x) for x in parsed] return None - if isinstance(parsed, list): - return [float(x) for x in parsed] + except (ValueError, TypeError): return None return None @@ -458,11 +461,12 @@ def _parse_row(self, row, include_embedding: bool = False) -> dict[str, Any]: } if include_embedding and len(row) > 5: result["metadata"]["embedding"] = row[5] - normalized = _normalize_embedding_value(result["metadata"].get("embedding")) - if normalized is not None: - result["metadata"]["embedding"] = normalized - elif "embedding" in result["metadata"]: - del result["metadata"]["embedding"] + if include_embedding: + normalized = _normalize_embedding_value(result["metadata"].get("embedding")) + if normalized is not None: + result["metadata"]["embedding"] = normalized + elif "embedding" in result["metadata"]: + del result["metadata"]["embedding"] return result @staticmethod diff --git a/tests/graph_dbs/test_postgres_embedding_parse.py b/tests/graph_dbs/test_postgres_embedding_parse.py index b66f8404e..14b608ea9 100644 --- a/tests/graph_dbs/test_postgres_embedding_parse.py +++ b/tests/graph_dbs/test_postgres_embedding_parse.py @@ -23,6 +23,11 @@ def test_normalize_embedding_value_rejects_invalid_string() -> None: assert _normalize_embedding_value("not-json") is None +def test_normalize_embedding_value_rejects_non_numeric_elements() -> None: + assert _normalize_embedding_value(["a", 0.5]) is None + assert _normalize_embedding_value('["a", 0.5]') is None + + def test_prepare_node_metadata_normalizes_string_embedding() -> None: metadata = _prepare_node_metadata({"embedding": "[0.25, 0.75]"}) assert metadata["embedding"] == [0.25, 0.75] @@ -54,3 +59,18 @@ def test_parse_row_normalizes_string_embedding_from_vector_column() -> None: parsed = db._parse_row(row, include_embedding=True) assert parsed["metadata"]["embedding"] == [0.25, 0.75] + + +def test_parse_row_preserves_props_embedding_when_not_requested() -> None: + db = _build_db() + row: tuple[Any, ...] = ( + "node-1", + "memory text", + json.dumps({"memory_type": "UserMemory", "embedding": "[0.1, 0.2]"}), + datetime(2026, 8, 22, 12, 0, 0), + datetime(2026, 8, 22, 12, 0, 0), + ) + + parsed = db._parse_row(row, include_embedding=False) + + assert parsed["metadata"]["embedding"] == "[0.1, 0.2]" From 26549913cb687f61692ac64eb546aef81545ee86 Mon Sep 17 00:00:00 2001 From: Kwizii Date: Sat, 22 Aug 2026 23:33:47 +0800 Subject: [PATCH 3/4] fix: add Postgres reorganizer/handler compatibility methods --- src/memos/graph_dbs/postgres.py | 155 +++++++++++++++++- .../test_postgres_reorganizer_compat.py | 82 +++++++++ 2 files changed, 230 insertions(+), 7 deletions(-) create mode 100644 tests/graph_dbs/test_postgres_reorganizer_compat.py diff --git a/src/memos/graph_dbs/postgres.py b/src/memos/graph_dbs/postgres.py index 8a4f07834..aaa2d4187 100644 --- a/src/memos/graph_dbs/postgres.py +++ b/src/memos/graph_dbs/postgres.py @@ -204,6 +204,49 @@ def _init_schema(self): finally: self._put_conn(conn) + def get_memory_count(self, memory_type: str, user_name: str | None = None) -> int: + """Count memory nodes by memory_type for a user.""" + user_name = user_name or self.user_name + conn = self._get_conn() + try: + with conn.cursor() as cur: + cur.execute( + f""" + SELECT COUNT(*) + FROM {self.schema}.memories + WHERE properties->>'memory_type' = %s + AND user_name = %s + """, + (memory_type, user_name), + ) + row = cur.fetchone() + return int(row[0]) if row else 0 + except Exception as e: + logger.error("[get_memory_count] Failed: %s", e) + return -1 + finally: + self._put_conn(conn) + + def node_not_exist(self, scope: str, user_name: str | None = None) -> bool: + """Return True when no activated nodes exist for the given memory_type scope.""" + user_name = user_name or self.user_name + conn = self._get_conn() + try: + with conn.cursor() as cur: + cur.execute( + f""" + SELECT 1 + FROM {self.schema}.memories + WHERE properties->>'memory_type' = %s + AND user_name = %s + LIMIT 1 + """, + (scope, user_name), + ) + return cur.fetchone() is None + finally: + self._put_conn(conn) + # ========================================================================= # Node Management # ========================================================================= @@ -704,23 +747,104 @@ def delete_edge( finally: self._put_conn(conn) - def edge_exists(self, source_id: str, target_id: str, type: str) -> bool: - """Check if edge exists.""" + def edge_exists( + self, + source_id: str, + target_id: str, + type: str = "ANY", + direction: str = "OUTGOING", + user_name: str | None = None, + ) -> bool: + """Check if an edge exists between two nodes.""" + user_name = user_name or self.user_name + if direction not in ("OUTGOING", "INCOMING", "ANY"): + raise ValueError( + f"Invalid direction: {direction}. Must be 'OUTGOING', 'INCOMING', or 'ANY'." + ) + + type_clause = "" if type == "ANY" else " AND e.edge_type = %s" + params: list[Any] = [user_name, user_name] + + if direction == "OUTGOING": + direction_clause = "e.source_id = %s AND e.target_id = %s" + params.extend([source_id, target_id]) + elif direction == "INCOMING": + direction_clause = "e.source_id = %s AND e.target_id = %s" + params.extend([target_id, source_id]) + else: + direction_clause = ( + "(e.source_id = %s AND e.target_id = %s) OR (e.source_id = %s AND e.target_id = %s)" + ) + params.extend([source_id, target_id, target_id, source_id]) + + if type != "ANY": + params.append(type) + conn = self._get_conn() try: with conn.cursor() as cur: cur.execute( f""" - SELECT 1 FROM {self.schema}.edges - WHERE source_id = %s AND target_id = %s AND edge_type = %s + SELECT 1 + FROM {self.schema}.edges e + JOIN {self.schema}.memories src ON src.id = e.source_id + JOIN {self.schema}.memories tgt ON tgt.id = e.target_id + WHERE src.user_name = %s + AND tgt.user_name = %s + AND ({direction_clause}) + {type_clause} LIMIT 1 """, - (source_id, target_id, type), + params, ) return cur.fetchone() is not None finally: self._put_conn(conn) + def get_edges( + self, id: str, type: str = "ANY", direction: str = "ANY", user_name: str | None = None + ) -> list[dict[str, str]]: + """Get edges connected to a node, with optional type and direction filter.""" + user_name = user_name or self.user_name + if direction not in ("OUTGOING", "INCOMING", "ANY"): + raise ValueError("Invalid direction. Must be 'OUTGOING', 'INCOMING', or 'ANY'.") + + type_clause = "" if type == "ANY" else " AND e.edge_type = %s" + params: list[Any] = [user_name, user_name, id] + + if direction == "OUTGOING": + node_clause = "e.source_id = %s" + elif direction == "INCOMING": + node_clause = "e.target_id = %s" + else: + node_clause = "(e.source_id = %s OR e.target_id = %s)" + params.append(id) + + if type != "ANY": + params.append(type) + + conn = self._get_conn() + try: + with conn.cursor() as cur: + cur.execute( + f""" + SELECT e.source_id, e.target_id, e.edge_type + FROM {self.schema}.edges e + JOIN {self.schema}.memories src ON src.id = e.source_id + JOIN {self.schema}.memories tgt ON tgt.id = e.target_id + WHERE src.user_name = %s + AND tgt.user_name = %s + AND {node_clause} + {type_clause} + """, + params, + ) + return [ + {"from": row[0], "to": row[1], "type": row[2]} for row in cur.fetchall() + ] + finally: + self._put_conn(conn) + # ========================================================================= # Graph Queries # ========================================================================= @@ -987,10 +1111,10 @@ def get_all_memory_items( self._put_conn(conn) def get_structure_optimization_candidates( - self, scope: str, include_embedding: bool = False + self, scope: str, include_embedding: bool = False, **kwargs ) -> list[dict]: """Find isolated nodes (no edges).""" - user_name = self.user_name + user_name = kwargs.get("user_name") or self.user_name conn = self._get_conn() try: with conn.cursor() as cur: @@ -1013,6 +1137,23 @@ def get_structure_optimization_candidates( finally: self._put_conn(conn) + def search_by_fulltext( + self, + query_words: list[str], + top_k: int = 10, + scope: str | None = None, + status: str | None = None, + threshold: float | None = None, + search_filter: dict | None = None, + user_name: str | None = None, + filter: dict | None = None, + knowledgebase_ids: list[str] | None = None, + tsquery_config: str | None = None, + **kwargs, + ) -> list[dict]: + """Stub for TreeTextMemory keyword recall; Postgres fulltext search is not implemented yet.""" + return [] + # ========================================================================= # Maintenance # ========================================================================= diff --git a/tests/graph_dbs/test_postgres_reorganizer_compat.py b/tests/graph_dbs/test_postgres_reorganizer_compat.py new file mode 100644 index 000000000..fef87b961 --- /dev/null +++ b/tests/graph_dbs/test_postgres_reorganizer_compat.py @@ -0,0 +1,82 @@ +"""Regression tests for PostgresGraphDB reorganizer/handler compatibility.""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +from memos.graph_dbs.postgres import PostgresGraphDB + + +def _build_db() -> PostgresGraphDB: + with ( + patch("memos.graph_dbs.postgres.require_python_package", lambda **kwargs: lambda fn: fn), + patch("psycopg2.pool.ThreadedConnectionPool", MagicMock()), + patch.object(PostgresGraphDB, "_init_schema", lambda self: None), + ): + config = MagicMock() + config.schema_name = "test_schema" + config.user_name = "user-1" + return PostgresGraphDB(config) + + +def _mock_cursor(fetchone=None, fetchall=None): + cursor = MagicMock() + cursor.fetchone.return_value = fetchone + cursor.fetchall.return_value = fetchall or [] + return cursor + + +def _mock_conn(cursor: MagicMock) -> MagicMock: + conn = MagicMock() + conn.cursor.return_value.__enter__.return_value = cursor + return conn + + +def test_node_not_exist_returns_true_when_no_nodes() -> None: + db = _build_db() + cursor = _mock_cursor(fetchone=None) + with patch.object(db, "_get_conn", return_value=_mock_conn(cursor)): + assert db.node_not_exist("LongTermMemory", user_name="user-1") is True + + +def test_node_not_exist_returns_false_when_nodes_exist() -> None: + db = _build_db() + cursor = _mock_cursor(fetchone=(1,)) + with patch.object(db, "_get_conn", return_value=_mock_conn(cursor)): + assert db.node_not_exist("LongTermMemory", user_name="user-1") is False + + +def test_get_memory_count_returns_count() -> None: + db = _build_db() + cursor = _mock_cursor(fetchone=(7,)) + with patch.object(db, "_get_conn", return_value=_mock_conn(cursor)): + assert db.get_memory_count("LongTermMemory", user_name="user-1") == 7 + + +def test_get_edges_outgoing() -> None: + db = _build_db() + cursor = _mock_cursor(fetchall=[("a", "b", "MERGED_TO")]) + with patch.object(db, "_get_conn", return_value=_mock_conn(cursor)): + edges = db.get_edges("a", type="ANY", direction="OUTGOING", user_name="user-1") + assert edges == [{"from": "a", "to": "b", "type": "MERGED_TO"}] + + +def test_edge_exists_any_direction() -> None: + db = _build_db() + cursor = _mock_cursor(fetchone=(1,)) + with patch.object(db, "_get_conn", return_value=_mock_conn(cursor)): + assert db.edge_exists("a", "b", "MERGED_TO", direction="ANY", user_name="user-1") is True + + +def test_get_structure_optimization_candidates_accepts_user_name_kwarg() -> None: + db = _build_db() + cursor = _mock_cursor(fetchall=[]) + with patch.object(db, "_get_conn", return_value=_mock_conn(cursor)): + result = db.get_structure_optimization_candidates("LongTermMemory", user_name="other-user") + assert result == [] + assert cursor.execute.call_args[0][1] == ("LongTermMemory", "other-user") + + +def test_search_by_fulltext_stub_returns_empty_list() -> None: + db = _build_db() + assert db.search_by_fulltext(["hello"], user_name="user-1") == [] From 560121e4d4903c9bdc3a19ba1900aa4d296008f8 Mon Sep 17 00:00:00 2001 From: Kwizii Date: Sat, 22 Aug 2026 23:42:23 +0800 Subject: [PATCH 4/4] fix: address OpenCodeReview findings for Postgres reorganizer PR --- src/memos/graph_dbs/postgres.py | 3 +-- .../test_postgres_embedding_parse.py | 19 ++++++++++++++----- .../test_postgres_reorganizer_compat.py | 4 ++-- 3 files changed, 17 insertions(+), 9 deletions(-) diff --git a/src/memos/graph_dbs/postgres.py b/src/memos/graph_dbs/postgres.py index aaa2d4187..a8fe7f1ea 100644 --- a/src/memos/graph_dbs/postgres.py +++ b/src/memos/graph_dbs/postgres.py @@ -504,11 +504,10 @@ def _parse_row(self, row, include_embedding: bool = False) -> dict[str, Any]: } if include_embedding and len(row) > 5: result["metadata"]["embedding"] = row[5] - if include_embedding: normalized = _normalize_embedding_value(result["metadata"].get("embedding")) if normalized is not None: result["metadata"]["embedding"] = normalized - elif "embedding" in result["metadata"]: + else: del result["metadata"]["embedding"] return result diff --git a/tests/graph_dbs/test_postgres_embedding_parse.py b/tests/graph_dbs/test_postgres_embedding_parse.py index 14b608ea9..7f671f34b 100644 --- a/tests/graph_dbs/test_postgres_embedding_parse.py +++ b/tests/graph_dbs/test_postgres_embedding_parse.py @@ -8,7 +8,13 @@ from typing import Any from unittest.mock import MagicMock, patch -from memos.graph_dbs.postgres import PostgresGraphDB, _normalize_embedding_value, _prepare_node_metadata +import pytest + +from memos.graph_dbs.postgres import ( + PostgresGraphDB, + _normalize_embedding_value, + _prepare_node_metadata, +) def test_normalize_embedding_value_parses_json_string() -> None: @@ -23,9 +29,12 @@ def test_normalize_embedding_value_rejects_invalid_string() -> None: assert _normalize_embedding_value("not-json") is None -def test_normalize_embedding_value_rejects_non_numeric_elements() -> None: - assert _normalize_embedding_value(["a", 0.5]) is None - assert _normalize_embedding_value('["a", 0.5]') is None +@pytest.mark.parametrize( + "embedding", + [["a", 0.5], '["a", 0.5]'], +) +def test_normalize_embedding_value_rejects_non_numeric_elements(embedding) -> None: + assert _normalize_embedding_value(embedding) is None def test_prepare_node_metadata_normalizes_string_embedding() -> None: @@ -35,7 +44,7 @@ def test_prepare_node_metadata_normalizes_string_embedding() -> None: def _build_db() -> PostgresGraphDB: with ( - patch("memos.graph_dbs.postgres.require_python_package", lambda **kwargs: lambda fn: fn), + patch("memos.graph_dbs.postgres.require_python_package", lambda *args, **kwargs: lambda fn: fn), patch("psycopg2.pool.ThreadedConnectionPool", MagicMock()), patch.object(PostgresGraphDB, "_init_schema", lambda self: None), ): diff --git a/tests/graph_dbs/test_postgres_reorganizer_compat.py b/tests/graph_dbs/test_postgres_reorganizer_compat.py index fef87b961..a77b23c84 100644 --- a/tests/graph_dbs/test_postgres_reorganizer_compat.py +++ b/tests/graph_dbs/test_postgres_reorganizer_compat.py @@ -9,7 +9,7 @@ def _build_db() -> PostgresGraphDB: with ( - patch("memos.graph_dbs.postgres.require_python_package", lambda **kwargs: lambda fn: fn), + patch("memos.graph_dbs.postgres.require_python_package", lambda *args, **kwargs: lambda fn: fn), patch("psycopg2.pool.ThreadedConnectionPool", MagicMock()), patch.object(PostgresGraphDB, "_init_schema", lambda self: None), ): @@ -22,7 +22,7 @@ def _build_db() -> PostgresGraphDB: def _mock_cursor(fetchone=None, fetchall=None): cursor = MagicMock() cursor.fetchone.return_value = fetchone - cursor.fetchall.return_value = fetchall or [] + cursor.fetchall.return_value = [] if fetchall is None else fetchall return cursor