fix(vector_stores): normalize scores to similarity (higher = better) across all backends (#5391)

This commit is contained in:
Kartik
2026-06-05 19:13:26 +05:30
committed by GitHub
parent b00a1a1065
commit 7ac8ab154b
20 changed files with 970 additions and 50 deletions
+2 -2
View File
@@ -45,7 +45,7 @@ def test_search_vectors(chromadb_instance, mock_chromadb_client):
assert len(results) == 2
assert results[0].id == "id1"
assert results[0].score == 0.1
assert results[0].score == pytest.approx(1.0 / 1.1)
assert results[0].payload == {"name": "vector1"}
@@ -150,7 +150,7 @@ def test_get_vector(chromadb_instance):
chromadb_instance.collection.get.assert_called_once_with(ids=["id1"])
assert result.id == "id1"
assert result.score == 0.1
assert result.score == pytest.approx(1.0 / 1.1)
assert result.payload == {"name": "vector1"}
+433
View File
@@ -0,0 +1,433 @@
"""
Level 2: End-to-end tests for Memory.search(threshold=...) across vector stores.
Tests the full pipeline: Memory.add() -> Memory.search(threshold=X) -> verify
that threshold filtering works correctly now that scores are similarity
(higher = better).
Before the fix, threshold filtering was inverted — good matches were dropped
and bad matches passed through. These tests verify the fix works end-to-end
through the Memory class, not just at the vector store layer.
In-memory providers (FAISS, ChromaDB) always run. External providers
(PGVector, Redis, Milvus, etc.) are skipped unless the service is reachable.
Set OPENAI_API_KEY env var or configure an alternative LLM/embedder to use
the full Memory pipeline; otherwise tests fall back to direct vector store
operations with synthetic embeddings.
Refs: https://github.com/mem0ai/mem0/issues/4453
"""
import os
import uuid
import numpy as np
import pytest
DIMS = 128
def _tcp_reachable(host, port, timeout=2):
import socket
try:
with socket.create_connection((host, port), timeout=timeout):
return True
except OSError:
return False
def _make_vectors():
"""Create 5 vectors with known similarity spread to a query."""
np.random.seed(42)
query = np.random.randn(DIMS).astype(np.float32)
query = query / np.linalg.norm(query)
vecs = []
for scale in [0.05, 0.15, 0.4, 0.8, 1.5]:
v = query + np.random.randn(DIMS).astype(np.float32) * scale
v = v / np.linalg.norm(v)
vecs.append(v.tolist())
return query.tolist(), vecs
def _run_threshold_test(store, query, doc_vectors, payloads, ids):
"""
Core threshold test logic shared across all providers.
Uses direct vector store API (not Memory class) so we can control
the exact vectors and test threshold behavior precisely.
"""
store.insert(vectors=doc_vectors, payloads=payloads, ids=ids)
# Step 1: Search without threshold (baseline)
results = store.search(query="", vectors=query, top_k=5)
assert len(results) > 0, "Baseline search returned no results"
scores = [r.score for r in results]
# Verify scores are similarity (higher = better)
assert scores[0] >= scores[-1], (
f"Top result should have highest score: first={scores[0]}, last={scores[-1]}"
)
# Step 2: Verify all scores are non-negative (similarity, not raw distance)
assert all(s >= 0 for s in scores if s is not None), (
f"All scores must be non-negative (not raw distances): {scores}"
)
# Step 3: Simulate threshold filtering as Memory.search() does it
# The check in mem0/memory/main.py is: if threshold is None or mem.score >= threshold
mid_threshold = (scores[0] + scores[-1]) / 2 if len(scores) >= 2 else scores[0] * 0.5
filtered = [r for r in results if r.score >= mid_threshold]
assert len(filtered) < len(results), (
f"Mid threshold {mid_threshold:.4f} should filter some results. "
f"Scores: {scores}"
)
assert len(filtered) > 0, (
f"Mid threshold {mid_threshold:.4f} should keep some results. "
f"Scores: {scores}"
)
# All filtered results must have score >= threshold
for r in filtered:
assert r.score >= mid_threshold, (
f"Score {r.score:.4f} below threshold {mid_threshold:.4f}"
)
# Step 4: Very high threshold should return 0 or very few results
high_threshold = 0.99
high_filtered = [r for r in results if r.score >= high_threshold]
assert len(high_filtered) < len(results), (
f"Threshold 0.99 should filter most results. Scores: {scores}"
)
return scores
# ---------------------------------------------------------------------------
# In-memory stores (always available)
# ---------------------------------------------------------------------------
class TestChromaDBThreshold:
def test_threshold_filtering(self, tmp_path):
from mem0.vector_stores.chroma import ChromaDB
store = ChromaDB(collection_name="test_threshold", path=str(tmp_path / "chroma"))
query, doc_vectors = _make_vectors()
payloads = [{"label": f"doc_{i}"} for i in range(5)]
ids = [f"id_{i}" for i in range(5)]
scores = _run_threshold_test(store, query, doc_vectors, payloads, ids)
assert all(0 < s <= 1.0 for s in scores), f"ChromaDB scores in (0,1]: {scores}"
store.delete_col()
def test_threshold_direction_not_inverted(self, tmp_path):
"""Regression test: before the fix, threshold filtering was inverted."""
from mem0.vector_stores.chroma import ChromaDB
store = ChromaDB(collection_name="test_inversion", path=str(tmp_path / "chroma2"))
query, doc_vectors = _make_vectors()
payloads = [{"label": f"doc_{i}"} for i in range(5)]
ids = [f"id_{i}" for i in range(5)]
store.insert(vectors=doc_vectors, payloads=payloads, ids=ids)
results = store.search(query="", vectors=query, top_k=5)
scores = [r.score for r in results]
# The bug was: all scores collapsed to 1.0 because raw L2 distances
# > 1.0 were capped. Verify scores are NOT all identical.
unique_scores = set(round(s, 6) for s in scores)
assert len(unique_scores) > 1, (
f"Scores should not all be identical (bug symptom): {scores}"
)
# The closest doc should score strictly higher than the farthest
assert scores[0] > scores[-1], (
f"Closest doc must score higher than farthest: {scores}"
)
store.delete_col()
class TestFAISSEuclideanThreshold:
def test_threshold_filtering(self, tmp_path):
from mem0.vector_stores.faiss import FAISS
store = FAISS(
collection_name="test_threshold",
path=str(tmp_path / "faiss"),
distance_strategy="euclidean",
embedding_model_dims=DIMS,
)
query, doc_vectors = _make_vectors()
payloads = [{"label": f"doc_{i}"} for i in range(5)]
ids = [f"id_{i}" for i in range(5)]
scores = _run_threshold_test(store, query, doc_vectors, payloads, ids)
assert all(0 < s <= 1.0 for s in scores), f"FAISS euclidean scores in (0,1]: {scores}"
class TestFAISSCosineThreshold:
def test_threshold_filtering(self, tmp_path):
from mem0.vector_stores.faiss import FAISS
store = FAISS(
collection_name="test_threshold",
path=str(tmp_path / "faiss_cos"),
distance_strategy="cosine",
embedding_model_dims=DIMS,
)
query, doc_vectors = _make_vectors()
payloads = [{"label": f"doc_{i}"} for i in range(5)]
ids = [f"id_{i}" for i in range(5)]
store.insert(vectors=doc_vectors, payloads=payloads, ids=ids)
results = store.search(query="", vectors=query, top_k=5)
scores = [r.score for r in results]
assert scores[0] >= scores[-1], f"Descending order: {scores}"
# ---------------------------------------------------------------------------
# External stores
# ---------------------------------------------------------------------------
PGVECTOR_HOST = os.environ.get("PGVECTOR_HOST", "localhost")
PGVECTOR_PORT = int(os.environ.get("PGVECTOR_PORT", "5432"))
PGVECTOR_USER = os.environ.get("PGVECTOR_USER", "mem0")
PGVECTOR_PASS = os.environ.get("PGVECTOR_PASSWORD", "mem0test")
PGVECTOR_DB = os.environ.get("PGVECTOR_DB", "mem0_test")
def _pgvector_reachable():
try:
import psycopg
conn = psycopg.connect(
host=PGVECTOR_HOST, port=PGVECTOR_PORT,
user=PGVECTOR_USER, password=PGVECTOR_PASS, dbname=PGVECTOR_DB,
connect_timeout=3,
)
conn.close()
return True
except Exception:
return False
@pytest.mark.skipif(
not _pgvector_reachable(),
reason=f"pgvector not reachable at {PGVECTOR_HOST}:{PGVECTOR_PORT} with user {PGVECTOR_USER}",
)
class TestPGVectorThreshold:
def test_threshold_filtering(self):
from mem0.vector_stores.pgvector import PGVector
collection = f"test_thr_{uuid.uuid4().hex[:8]}"
store = PGVector(
collection_name=collection,
embedding_model_dims=DIMS,
host=PGVECTOR_HOST,
port=PGVECTOR_PORT,
user=os.environ.get("PGVECTOR_USER", "mem0"),
password=os.environ.get("PGVECTOR_PASSWORD", "mem0test"),
dbname=os.environ.get("PGVECTOR_DB", "mem0_test"),
diskann=False,
hnsw=True,
)
query, doc_vectors = _make_vectors()
payloads = [{"label": f"doc_{i}"} for i in range(5)]
ids = [str(uuid.uuid4()) for _ in range(5)]
scores = _run_threshold_test(store, query, doc_vectors, payloads, ids)
assert all(0 <= s <= 1.0 for s in scores), f"PGVector scores in [0,1]: {scores}"
store.delete_col()
REDIS_HOST = os.environ.get("REDIS_HOST", "localhost")
REDIS_PORT = int(os.environ.get("REDIS_PORT", "6379"))
@pytest.mark.skipif(
not _tcp_reachable(REDIS_HOST, REDIS_PORT),
reason=f"Redis not reachable at {REDIS_HOST}:{REDIS_PORT}",
)
class TestRedisThreshold:
def test_threshold_filtering(self):
from datetime import datetime, timezone
from mem0.vector_stores.redis import RedisDB
collection = f"test_thr_{uuid.uuid4().hex[:8]}"
store = RedisDB(
collection_name=collection,
embedding_model_dims=DIMS,
redis_url=f"redis://{REDIS_HOST}:{REDIS_PORT}",
)
query, doc_vectors = _make_vectors()
now = datetime.now(timezone.utc).isoformat(timespec="microseconds")
payloads = [
{"hash": f"h{i}", "data": f"doc_{i} memory", "created_at": now, "user_id": "test", "label": f"doc_{i}"}
for i in range(5)
]
ids = [str(uuid.uuid4()) for _ in range(5)]
store.insert(vectors=doc_vectors, payloads=payloads, ids=ids)
results = store.search(query="", vectors=query, top_k=5, filters={"user_id": "test"})
scores = [r.score for r in results]
assert all(0 <= s <= 1.0 for s in scores), f"Redis scores in [0,1]: {scores}"
assert scores[0] >= scores[-1], f"Descending order: {scores}"
mid = (scores[0] + scores[-1]) / 2
filtered = [r for r in results if r.score >= mid]
assert 0 < len(filtered) < len(results), f"Threshold {mid} should filter: {scores}"
store.delete_col()
VALKEY_HOST = os.environ.get("VALKEY_HOST", "localhost")
VALKEY_PORT = int(os.environ.get("VALKEY_PORT", "6380"))
@pytest.mark.skipif(
not _tcp_reachable(VALKEY_HOST, VALKEY_PORT),
reason=f"Valkey not reachable at {VALKEY_HOST}:{VALKEY_PORT}",
)
class TestValkeyThreshold:
def test_threshold_filtering(self):
from datetime import datetime, timezone
from mem0.vector_stores.valkey import ValkeyDB
collection = f"test_thr_{uuid.uuid4().hex[:8]}"
store = ValkeyDB(
collection_name=collection,
embedding_model_dims=DIMS,
valkey_url=f"valkey://{VALKEY_HOST}:{VALKEY_PORT}",
)
query, doc_vectors = _make_vectors()
now = datetime.now(timezone.utc).isoformat(timespec="microseconds")
payloads = [
{"hash": f"h{i}", "data": f"doc_{i} memory", "created_at": now, "user_id": "test", "label": f"doc_{i}"}
for i in range(5)
]
ids = [str(uuid.uuid4()) for _ in range(5)]
store.insert(vectors=doc_vectors, payloads=payloads, ids=ids)
results = store.search(query="", vectors=query, top_k=5, filters={"user_id": "test"})
scores = [r.score for r in results]
assert all(0 <= s <= 1.0 for s in scores), f"Valkey scores in [0,1]: {scores}"
assert scores[0] >= scores[-1], f"Descending order: {scores}"
store.delete_col()
MILVUS_HOST = os.environ.get("MILVUS_HOST", "localhost")
MILVUS_PORT = int(os.environ.get("MILVUS_PORT", "19530"))
@pytest.mark.skipif(
not _tcp_reachable(MILVUS_HOST, MILVUS_PORT),
reason=f"Milvus not reachable at {MILVUS_HOST}:{MILVUS_PORT}",
)
class TestMilvusL2Threshold:
def test_threshold_filtering(self):
from mem0.vector_stores.milvus import MilvusDB
collection = f"test_l2_{uuid.uuid4().hex[:8]}"
store = MilvusDB(
collection_name=collection,
embedding_model_dims=DIMS,
url=f"http://{MILVUS_HOST}:{MILVUS_PORT}",
token="",
db_name="",
metric_type="L2",
)
query, doc_vectors = _make_vectors()
payloads = [{"label": f"doc_{i}"} for i in range(5)]
ids = [str(uuid.uuid4()) for _ in range(5)]
store.insert(ids=ids, vectors=doc_vectors, payloads=payloads)
scores = _run_threshold_test(store, query, doc_vectors[:3], payloads[:3], [str(uuid.uuid4()) for _ in range(3)])
# L2 scores via 1/(1+d) should be in (0, 1]
assert all(0 < s <= 1.0 for s in scores), f"Milvus L2 scores in (0,1]: {scores}"
store.delete_col()
@pytest.mark.skipif(
not _tcp_reachable(MILVUS_HOST, MILVUS_PORT),
reason=f"Milvus not reachable at {MILVUS_HOST}:{MILVUS_PORT}",
)
class TestMilvusCosineThreshold:
def test_threshold_filtering(self):
from mem0.vector_stores.milvus import MilvusDB
collection = f"test_cos_{uuid.uuid4().hex[:8]}"
store = MilvusDB(
collection_name=collection,
embedding_model_dims=DIMS,
url=f"http://{MILVUS_HOST}:{MILVUS_PORT}",
token="",
db_name="",
metric_type="COSINE",
)
query, doc_vectors = _make_vectors()
payloads = [{"label": f"doc_{i}"} for i in range(5)]
ids = [str(uuid.uuid4()) for _ in range(5)]
store.insert(ids=ids, vectors=doc_vectors, payloads=payloads)
results = store.search(query="", vectors=query, top_k=5)
scores = [r.score for r in results]
assert scores[0] >= scores[-1], f"Descending: {scores}"
store.delete_col()
SUPABASE_CONN = os.environ.get("SUPABASE_CONN_STRING", "")
@pytest.mark.skipif(not SUPABASE_CONN, reason="SUPABASE_CONN_STRING not set")
class TestSupabaseThreshold:
def test_threshold_filtering(self):
from mem0.vector_stores.supabase import Supabase
collection = f"test_thr_{uuid.uuid4().hex[:8]}"
store = Supabase(
connection_string=SUPABASE_CONN,
collection_name=collection,
embedding_model_dims=DIMS,
)
query, doc_vectors = _make_vectors()
payloads = [{"label": f"doc_{i}"} for i in range(5)]
ids = [f"id_{i}" for i in range(5)]
scores = _run_threshold_test(store, query, doc_vectors, payloads, ids)
assert all(0 <= s <= 1.0 for s in scores), f"Supabase scores in [0,1]: {scores}"
store.delete_col()
S3_BUCKET = os.environ.get("S3_VECTORS_BUCKET", "")
@pytest.mark.skipif(not S3_BUCKET, reason="S3_VECTORS_BUCKET not set")
class TestS3VectorsThreshold:
def test_threshold_filtering(self):
from mem0.vector_stores.s3_vectors import S3Vectors
collection = f"testthr{uuid.uuid4().hex[:8]}"
region = os.environ.get("S3_VECTORS_REGION", "us-east-1")
store = S3Vectors(
vector_bucket_name=S3_BUCKET,
collection_name=collection,
embedding_model_dims=DIMS,
distance_metric="cosine",
region_name=region,
)
query, doc_vectors = _make_vectors()
payloads = [{"label": f"doc_{i}"} for i in range(5)]
ids = [f"id_{i}" for i in range(5)]
scores = _run_threshold_test(store, query, doc_vectors, payloads, ids)
assert all(0 <= s <= 1.0 for s in scores), f"S3 scores in [0,1]: {scores}"
store.delete_col()
+17 -13
View File
@@ -4,7 +4,11 @@ import unittest
import uuid
from unittest.mock import MagicMock, patch
from mem0.vector_stores.pgvector import PGVector, _build_filter_conditions, _with_sslmode
from mem0.vector_stores.pgvector import (
PGVector,
_build_filter_conditions,
_with_sslmode,
)
class TestPGVector(unittest.TestCase):
@@ -450,9 +454,9 @@ class TestPGVector(unittest.TestCase):
# Verify results
self.assertEqual(len(results), 2)
self.assertEqual(results[0].id, self.test_ids[0])
self.assertEqual(results[0].score, 0.1)
self.assertEqual(results[0].score, 0.9)
self.assertEqual(results[1].id, self.test_ids[1])
self.assertEqual(results[1].score, 0.2)
self.assertEqual(results[1].score, 0.8)
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
@patch('mem0.vector_stores.pgvector.ConnectionPool')
@@ -499,9 +503,9 @@ class TestPGVector(unittest.TestCase):
# Verify results
self.assertEqual(len(results), 2)
self.assertEqual(results[0].id, self.test_ids[0])
self.assertEqual(results[0].score, 0.1)
self.assertEqual(results[0].score, 0.9)
self.assertEqual(results[1].id, self.test_ids[1])
self.assertEqual(results[1].score, 0.2)
self.assertEqual(results[1].score, 0.8)
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
@patch('mem0.vector_stores.pgvector.ConnectionPool')
@@ -1145,7 +1149,7 @@ class TestPGVector(unittest.TestCase):
# Verify results
self.assertEqual(len(results), 1)
self.assertEqual(results[0].id, self.test_ids[0])
self.assertEqual(results[0].score, 0.1)
self.assertEqual(results[0].score, 0.9)
self.assertEqual(results[0].payload["user_id"], "alice")
self.assertEqual(results[0].payload["agent_id"], "agent1")
self.assertEqual(results[0].payload["run_id"], "run1")
@@ -1195,7 +1199,7 @@ class TestPGVector(unittest.TestCase):
# Verify results
self.assertEqual(len(results), 1)
self.assertEqual(results[0].id, self.test_ids[0])
self.assertEqual(results[0].score, 0.1)
self.assertEqual(results[0].score, 0.9)
self.assertEqual(results[0].payload["user_id"], "alice")
self.assertEqual(results[0].payload["agent_id"], "agent1")
self.assertEqual(results[0].payload["run_id"], "run1")
@@ -1245,7 +1249,7 @@ class TestPGVector(unittest.TestCase):
# Verify results
self.assertEqual(len(results), 1)
self.assertEqual(results[0].id, self.test_ids[0])
self.assertEqual(results[0].score, 0.1)
self.assertEqual(results[0].score, 0.9)
self.assertEqual(results[0].payload["user_id"], "alice")
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
@@ -1293,7 +1297,7 @@ class TestPGVector(unittest.TestCase):
# Verify results
self.assertEqual(len(results), 1)
self.assertEqual(results[0].id, self.test_ids[0])
self.assertEqual(results[0].score, 0.1)
self.assertEqual(results[0].score, 0.9)
self.assertEqual(results[0].payload["user_id"], "alice")
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
@@ -1341,9 +1345,9 @@ class TestPGVector(unittest.TestCase):
# Verify results
self.assertEqual(len(results), 2)
self.assertEqual(results[0].id, self.test_ids[0])
self.assertEqual(results[0].score, 0.1)
self.assertEqual(results[0].score, 0.9)
self.assertEqual(results[1].id, self.test_ids[1])
self.assertEqual(results[1].score, 0.2)
self.assertEqual(results[1].score, 0.8)
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
@patch('mem0.vector_stores.pgvector.ConnectionPool')
@@ -1390,9 +1394,9 @@ class TestPGVector(unittest.TestCase):
# Verify results
self.assertEqual(len(results), 2)
self.assertEqual(results[0].id, self.test_ids[0])
self.assertEqual(results[0].score, 0.1)
self.assertEqual(results[0].score, 0.9)
self.assertEqual(results[1].id, self.test_ids[1])
self.assertEqual(results[1].score, 0.2)
self.assertEqual(results[1].score, 0.8)
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
@patch('mem0.vector_stores.pgvector.ConnectionPool')
+1 -1
View File
@@ -157,7 +157,7 @@ def test_search(mock_boto_client):
mock_boto_client.query_vectors.assert_called_once()
assert len(results) == 1
assert results[0].id == "id1"
assert results[0].score == 0.9
assert results[0].score == pytest.approx(0.1)
def test_get(mock_boto_client):
@@ -0,0 +1,459 @@
"""
Level 1: Vector store layer tests for score normalization.
Verifies that all vector stores return similarity scores (higher = better)
after the distance-to-similarity conversion fix. Each test directly inserts
vectors with known similarity ordering, then asserts scores are:
- Non-negative
- In [0, 1] (for normalized stores)
- Descending (close > mid > far)
- Closest vector ranks first
- Threshold filtering works (score >= threshold keeps correct results)
Providers that need external services are skipped unless the service is
reachable. FAISS and ChromaDB run in-memory and always execute.
Refs: https://github.com/mem0ai/mem0/issues/4453
"""
import os
import uuid
import numpy as np
import pytest
DIMS = 128
@pytest.fixture(scope="module")
def vectors():
"""Create a query vector and 3 doc vectors with known similarity ordering."""
np.random.seed(42)
query = np.random.randn(DIMS).astype(np.float32)
query = query / np.linalg.norm(query)
close = query + np.random.randn(DIMS).astype(np.float32) * 0.1
close = close / np.linalg.norm(close)
mid = query + np.random.randn(DIMS).astype(np.float32) * 0.5
mid = mid / np.linalg.norm(mid)
far = np.random.randn(DIMS).astype(np.float32)
far = far / np.linalg.norm(far)
return {
"query": query.tolist(),
"docs": [close.tolist(), mid.tolist(), far.tolist()],
"payloads": [{"label": "close"}, {"label": "mid"}, {"label": "far"}],
"ids": ["close", "mid", "far"],
}
def _assert_similarity_scores(results, *, allow_negative=False):
"""Common assertions for normalized similarity scores."""
scores = [r.score for r in results]
labels = [r.payload.get("label", r.payload.get("data", "?").split()[0]) for r in results]
assert len(results) == 3, f"Expected 3 results, got {len(results)}"
if not allow_negative:
assert all(s >= 0 for s in scores), f"Scores must be non-negative: {scores}"
assert all(s <= 1.0 for s in scores), f"Scores must be <= 1.0: {scores}"
assert scores[0] >= scores[1] >= scores[2], (
f"Scores must be descending (higher=better): {list(zip(labels, scores))}"
)
assert labels[0] == "close", f"Closest vector must rank first, got: {labels[0]}"
# Threshold filtering: using mid score should keep at least close and mid
threshold = scores[1]
filtered = [r for r in results if r.score >= threshold]
assert len(filtered) >= 2, (
f"Threshold {threshold} should keep >= 2 results, got {len(filtered)}"
)
# ---------------------------------------------------------------------------
# In-memory stores (always available)
# ---------------------------------------------------------------------------
class TestChromaDB:
"""ChromaDB uses L2 distance. Conversion: score = 1 / (1 + distance)."""
@pytest.fixture(autouse=True)
def setup(self, vectors, tmp_path):
from mem0.vector_stores.chroma import ChromaDB
self.store = ChromaDB(collection_name="test_norm", path=str(tmp_path / "chroma"))
self.store.insert(
vectors=vectors["docs"],
payloads=vectors["payloads"],
ids=vectors["ids"],
)
self.query = vectors["query"]
yield
self.store.delete_col()
def test_scores_are_similarity(self):
results = self.store.search(query="", vectors=self.query, top_k=3)
_assert_similarity_scores(results)
def test_score_formula(self):
"""Verify the exact L2-to-similarity conversion."""
results = self.store.search(query="", vectors=self.query, top_k=3)
for r in results:
assert r.score > 0.0
assert r.score <= 1.0
class TestFAISSEuclidean:
"""FAISS euclidean uses L2 distance. Conversion: score = 1 / (1 + distance)."""
@pytest.fixture(autouse=True)
def setup(self, vectors, tmp_path):
from mem0.vector_stores.faiss import FAISS
self.store = FAISS(
collection_name="test_norm",
path=str(tmp_path / "faiss_euc"),
distance_strategy="euclidean",
embedding_model_dims=DIMS,
)
self.store.insert(
vectors=vectors["docs"],
payloads=vectors["payloads"],
ids=vectors["ids"],
)
self.query = vectors["query"]
yield
def test_scores_are_similarity(self):
results = self.store.search(query="", vectors=self.query, top_k=3)
_assert_similarity_scores(results)
class TestFAISSCosine:
"""FAISS cosine uses inner product (higher = better). No conversion needed."""
@pytest.fixture(autouse=True)
def setup(self, vectors, tmp_path):
from mem0.vector_stores.faiss import FAISS
self.store = FAISS(
collection_name="test_norm",
path=str(tmp_path / "faiss_cos"),
distance_strategy="cosine",
embedding_model_dims=DIMS,
)
self.store.insert(
vectors=vectors["docs"],
payloads=vectors["payloads"],
ids=vectors["ids"],
)
self.query = vectors["query"]
yield
def test_scores_are_descending(self):
results = self.store.search(query="", vectors=self.query, top_k=3)
_assert_similarity_scores(results, allow_negative=True)
# ---------------------------------------------------------------------------
# External stores (skipped if service unavailable)
# ---------------------------------------------------------------------------
def _tcp_reachable(host, port, timeout=2):
import socket
try:
with socket.create_connection((host, port), timeout=timeout):
return True
except OSError:
return False
# --- PGVector ---
PGVECTOR_HOST = os.environ.get("PGVECTOR_HOST", "localhost")
PGVECTOR_PORT = int(os.environ.get("PGVECTOR_PORT", "5432"))
PGVECTOR_USER = os.environ.get("PGVECTOR_USER", "mem0")
PGVECTOR_PASS = os.environ.get("PGVECTOR_PASSWORD", "mem0test")
PGVECTOR_DB = os.environ.get("PGVECTOR_DB", "mem0_test")
def _pgvector_reachable():
try:
import psycopg
conn = psycopg.connect(
host=PGVECTOR_HOST, port=PGVECTOR_PORT,
user=PGVECTOR_USER, password=PGVECTOR_PASS, dbname=PGVECTOR_DB,
connect_timeout=3,
)
conn.close()
return True
except Exception:
return False
@pytest.mark.skipif(
not _pgvector_reachable(),
reason=f"pgvector not reachable at {PGVECTOR_HOST}:{PGVECTOR_PORT} with user {PGVECTOR_USER}",
)
class TestPGVector:
"""PGVector uses cosine distance [0,2]. Conversion: score = max(0, 1 - dist)."""
@pytest.fixture(autouse=True)
def setup(self, vectors):
from mem0.vector_stores.pgvector import PGVector
self.collection = f"test_norm_{uuid.uuid4().hex[:8]}"
self.store = PGVector(
collection_name=self.collection,
embedding_model_dims=DIMS,
host=PGVECTOR_HOST,
port=PGVECTOR_PORT,
user=PGVECTOR_USER,
password=PGVECTOR_PASS,
dbname=PGVECTOR_DB,
diskann=False,
hnsw=True,
)
ids = [str(uuid.uuid4()) for _ in range(3)]
self.store.insert(vectors=vectors["docs"], payloads=vectors["payloads"], ids=ids)
self.query = vectors["query"]
yield
self.store.delete_col()
def test_scores_are_similarity(self):
results = self.store.search(query="", vectors=self.query, top_k=3)
_assert_similarity_scores(results)
# --- Redis ---
REDIS_HOST = os.environ.get("REDIS_HOST", "localhost")
REDIS_PORT = int(os.environ.get("REDIS_PORT", "6379"))
@pytest.mark.skipif(
not _tcp_reachable(REDIS_HOST, REDIS_PORT),
reason=f"Redis not reachable at {REDIS_HOST}:{REDIS_PORT}",
)
class TestRedis:
"""Redis returns cosine distance [0,2]. Conversion: score = max(0, 1 - dist)."""
@pytest.fixture(autouse=True)
def setup(self, vectors):
from datetime import datetime, timezone
from mem0.vector_stores.redis import RedisDB
self.collection = f"test_norm_{uuid.uuid4().hex[:8]}"
self.store = RedisDB(
collection_name=self.collection,
embedding_model_dims=DIMS,
redis_url=f"redis://{REDIS_HOST}:{REDIS_PORT}",
)
now = datetime.now(timezone.utc).isoformat(timespec="microseconds")
payloads = [
{"hash": f"h{i}", "data": f"{v['label']} memory", "created_at": now, "user_id": "test", **v}
for i, v in enumerate(vectors["payloads"])
]
ids = [str(uuid.uuid4()) for _ in range(3)]
self.store.insert(vectors=vectors["docs"], payloads=payloads, ids=ids)
self.query = vectors["query"]
yield
self.store.delete_col()
def test_scores_are_similarity(self):
results = self.store.search(
query="", vectors=self.query, top_k=3, filters={"user_id": "test"}
)
scores = [r.score for r in results]
assert all(s >= 0 for s in scores), f"Scores must be non-negative: {scores}"
assert all(s <= 1.0 for s in scores), f"Scores must be <= 1.0: {scores}"
assert scores[0] >= scores[1] >= scores[2], f"Scores must be descending: {scores}"
# --- Valkey ---
VALKEY_HOST = os.environ.get("VALKEY_HOST", "localhost")
VALKEY_PORT = int(os.environ.get("VALKEY_PORT", "6380"))
@pytest.mark.skipif(
not _tcp_reachable(VALKEY_HOST, VALKEY_PORT),
reason=f"Valkey not reachable at {VALKEY_HOST}:{VALKEY_PORT}",
)
class TestValkey:
"""Valkey returns cosine distance [0,2]. Conversion: score = max(0, 1 - dist)."""
@pytest.fixture(autouse=True)
def setup(self, vectors):
from datetime import datetime, timezone
from mem0.vector_stores.valkey import ValkeyDB
self.collection = f"test_norm_{uuid.uuid4().hex[:8]}"
self.store = ValkeyDB(
collection_name=self.collection,
embedding_model_dims=DIMS,
valkey_url=f"valkey://{VALKEY_HOST}:{VALKEY_PORT}",
)
now = datetime.now(timezone.utc).isoformat(timespec="microseconds")
payloads = [
{"hash": f"h{i}", "data": f"{v['label']} memory", "created_at": now, "user_id": "test", **v}
for i, v in enumerate(vectors["payloads"])
]
ids = [str(uuid.uuid4()) for _ in range(3)]
self.store.insert(vectors=vectors["docs"], payloads=payloads, ids=ids)
self.query = vectors["query"]
yield
self.store.delete_col()
def test_scores_are_similarity(self):
results = self.store.search(
query="", vectors=self.query, top_k=3, filters={"user_id": "test"}
)
scores = [r.score for r in results]
assert all(s >= 0 for s in scores), f"Scores must be non-negative: {scores}"
assert all(s <= 1.0 for s in scores), f"Scores must be <= 1.0: {scores}"
assert scores[0] >= scores[1] >= scores[2], f"Scores must be descending: {scores}"
# --- Milvus L2 ---
MILVUS_HOST = os.environ.get("MILVUS_HOST", "localhost")
MILVUS_PORT = int(os.environ.get("MILVUS_PORT", "19530"))
@pytest.mark.skipif(
not _tcp_reachable(MILVUS_HOST, MILVUS_PORT),
reason=f"Milvus not reachable at {MILVUS_HOST}:{MILVUS_PORT}",
)
class TestMilvusL2:
"""Milvus L2 returns L2 distance. Conversion: score = 1 / (1 + distance)."""
@pytest.fixture(autouse=True)
def setup(self, vectors):
from mem0.vector_stores.milvus import MilvusDB
self.collection = f"test_l2_{uuid.uuid4().hex[:8]}"
self.store = MilvusDB(
collection_name=self.collection,
embedding_model_dims=DIMS,
url=f"http://{MILVUS_HOST}:{MILVUS_PORT}",
token="",
db_name="",
metric_type="L2",
)
ids = [str(uuid.uuid4()) for _ in range(3)]
self.store.insert(ids=ids, vectors=vectors["docs"], payloads=vectors["payloads"])
self.query = vectors["query"]
yield
self.store.delete_col()
def test_scores_are_similarity(self):
results = self.store.search(query="", vectors=self.query, top_k=3)
_assert_similarity_scores(results)
@pytest.mark.skipif(
not _tcp_reachable(MILVUS_HOST, MILVUS_PORT),
reason=f"Milvus not reachable at {MILVUS_HOST}:{MILVUS_PORT}",
)
class TestMilvusCosine:
"""Milvus COSINE already returns similarity. No conversion needed."""
@pytest.fixture(autouse=True)
def setup(self, vectors):
from mem0.vector_stores.milvus import MilvusDB
self.collection = f"test_cos_{uuid.uuid4().hex[:8]}"
self.store = MilvusDB(
collection_name=self.collection,
embedding_model_dims=DIMS,
url=f"http://{MILVUS_HOST}:{MILVUS_PORT}",
token="",
db_name="",
metric_type="COSINE",
)
ids = [str(uuid.uuid4()) for _ in range(3)]
self.store.insert(ids=ids, vectors=vectors["docs"], payloads=vectors["payloads"])
self.query = vectors["query"]
yield
self.store.delete_col()
def test_scores_are_descending(self):
results = self.store.search(query="", vectors=self.query, top_k=3)
scores = [r.score for r in results]
labels = [r.payload["label"] for r in results]
assert scores[0] >= scores[1] >= scores[2], f"Descending: {list(zip(labels, scores))}"
assert labels[0] == "close"
# --- Supabase ---
SUPABASE_CONN = os.environ.get("SUPABASE_CONN_STRING", "")
@pytest.mark.skipif(not SUPABASE_CONN, reason="SUPABASE_CONN_STRING not set")
class TestSupabase:
"""Supabase (vecs) returns cosine distance [0,2]. Conversion: score = max(0, 1 - dist)."""
@pytest.fixture(autouse=True)
def setup(self, vectors):
from mem0.vector_stores.supabase import Supabase
self.collection = f"test_norm_{uuid.uuid4().hex[:8]}"
self.store = Supabase(
connection_string=SUPABASE_CONN,
collection_name=self.collection,
embedding_model_dims=DIMS,
)
self.store.insert(vectors=vectors["docs"], payloads=vectors["payloads"], ids=vectors["ids"])
self.query = vectors["query"]
yield
self.store.delete_col()
def test_scores_are_similarity(self):
results = self.store.search(query="", vectors=self.query, top_k=3)
_assert_similarity_scores(results)
# --- S3 Vectors ---
S3_BUCKET = os.environ.get("S3_VECTORS_BUCKET", "")
@pytest.mark.skipif(not S3_BUCKET, reason="S3_VECTORS_BUCKET not set")
class TestS3Vectors:
"""S3 Vectors returns cosine distance. Conversion: score = max(0, 1 - dist)."""
@pytest.fixture(autouse=True)
def setup(self, vectors):
from mem0.vector_stores.s3_vectors import S3Vectors
self.collection = f"testnorm{uuid.uuid4().hex[:8]}"
region = os.environ.get("S3_VECTORS_REGION", "us-east-1")
self.store = S3Vectors(
vector_bucket_name=S3_BUCKET,
collection_name=self.collection,
embedding_model_dims=DIMS,
distance_metric="cosine",
region_name=region,
)
self.store.insert(vectors=vectors["docs"], payloads=vectors["payloads"], ids=vectors["ids"])
self.query = vectors["query"]
yield
self.store.delete_col()
def test_scores_are_similarity(self):
results = self.store.search(query="", vectors=self.query, top_k=3)
_assert_similarity_scores(results)
+1 -1
View File
@@ -75,7 +75,7 @@ def test_search_vectors(supabase_instance, mock_collection):
assert len(results) == 2
assert results[0].id == "id1"
assert results[0].score == 0.9
assert results[0].score == pytest.approx(0.1)
assert results[0].payload == {"name": "vector1"}
@@ -112,7 +112,7 @@ def test_search_vectors(vector_store, mock_vertex_ai):
assert len(results) == 1
assert results[0].id == "test-id"
assert results[0].score == 0.1
assert results[0].score == 0.9
assert results[0].payload == {"user_id": "test_user"}