Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 34 additions & 4 deletions src/memos/graph_dbs/postgres.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,16 +26,40 @@
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
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
except (ValueError, TypeError):
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

Expand Down Expand Up @@ -437,6 +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]
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
Expand Down
76 changes: 76 additions & 0 deletions tests/graph_dbs/test_postgres_embedding_parse.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
"""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_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]


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]


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]"