From 05cd088b6e65c51f0149c6b660fcf3f4891e40be Mon Sep 17 00:00:00 2001 From: Kwizii Date: Sat, 22 Aug 2026 23:08:03 +0800 Subject: [PATCH 1/2] 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/2] 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]"