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
+1 -1
View File
@@ -444,7 +444,7 @@ export class RedisDB implements VectorStore {
return {
id: doc.value.memory_id,
payload: toCamelCase(resultPayload),
score: Number(doc.value.__vector_score) ?? 0,
score: Math.max(0, 1 - (Number(doc.value.__vector_score) ?? 0)),
};
});
} catch (error) {
+3 -6
View File
@@ -313,13 +313,10 @@ class AzureMySQL(VectorStoreBase):
for row in results:
vec = np.array(json.loads(row['vector']))
# Cosine similarity
similarity = np.dot(query_vec, vec) / (np.linalg.norm(query_vec) * np.linalg.norm(vec))
distance = 1 - similarity
scored_results.append((row['id'], distance, row['payload']))
similarity = float(np.dot(query_vec, vec) / (np.linalg.norm(query_vec) * np.linalg.norm(vec)))
scored_results.append((row['id'], similarity, row['payload']))
# Sort by distance and apply limit
scored_results.sort(key=lambda x: x[1])
scored_results.sort(key=lambda x: x[1], reverse=True)
scored_results = scored_results[:top_k]
return [
+9 -1
View File
@@ -14,7 +14,15 @@ class VectorStoreBase(ABC):
@abstractmethod
def search(self, query, vectors, top_k=5, filters=None):
"""Search for similar vectors."""
"""Search for similar vectors.
All implementations must return similarity scores where higher values
indicate greater similarity (range [0, 1] preferred). Implementations
using distance metrics must convert to similarity before returning:
- Cosine distance: score = max(0.0, 1.0 - distance)
- L2 distance: score = 1.0 / (1.0 + distance)
- Inner product: score = value (already higher = better)
"""
pass
@abstractmethod
+4 -7
View File
@@ -258,10 +258,8 @@ class CassandraDB(VectorStoreBase):
continue
vec = np.array(row.vector)
# Cosine similarity
similarity = np.dot(query_vec, vec) / (np.linalg.norm(query_vec) * np.linalg.norm(vec))
distance = 1 - similarity
similarity = float(np.dot(query_vec, vec) / (np.linalg.norm(query_vec) * np.linalg.norm(vec)))
# Apply filters if provided
if filters:
@@ -273,10 +271,9 @@ class CassandraDB(VectorStoreBase):
except json.JSONDecodeError:
continue
scored_results.append((row.id, distance, row.payload))
scored_results.append((row.id, similarity, row.payload))
# Sort by distance and apply limit
scored_results.sort(key=lambda x: x[1])
scored_results.sort(key=lambda x: x[1], reverse=True)
scored_results = scored_results[:top_k]
return [
+3 -1
View File
@@ -97,9 +97,11 @@ class ChromaDB(VectorStoreBase):
result = []
for i in range(max_length):
raw_distance = distances[i] if isinstance(distances, list) and distances and i < len(distances) else None
score = 1.0 / (1.0 + raw_distance) if raw_distance is not None else None
entry = OutputData(
id=ids[i] if isinstance(ids, list) and ids and i < len(ids) else None,
score=(distances[i] if isinstance(distances, list) and distances and i < len(distances) else None),
score=score,
payload=(metadatas[i] if isinstance(metadatas, list) and metadatas and i < len(metadatas) else None),
)
result.append(entry)
+5 -1
View File
@@ -279,7 +279,11 @@ class FAISS(VectorStoreBase):
payload_copy = payload.copy()
score = float(scores[i])
raw_score = float(scores[i])
if self.distance_strategy.lower() == "euclidean":
score = 1.0 / (1.0 + raw_score)
else:
score = raw_score
entry = OutputData(
id=vector_id,
score=score,
+16 -6
View File
@@ -11,7 +11,14 @@ try:
except ImportError:
raise ImportError("The 'pymilvus' library is required. Please install it using 'pip install pymilvus'.")
from pymilvus import CollectionSchema, DataType, FieldSchema, Function, FunctionType, MilvusClient
from pymilvus import (
CollectionSchema,
DataType,
FieldSchema,
Function,
FunctionType,
MilvusClient,
)
logger = logging.getLogger(__name__)
@@ -167,11 +174,14 @@ class MilvusDB(VectorStoreBase):
memory = []
for value in data:
uid, score, metadata = (
value.get("id"),
value.get("distance"),
value.get("entity", {}).get("metadata"),
)
uid = value.get("id")
raw_distance = value.get("distance")
metadata = value.get("entity", {}).get("metadata")
if raw_distance is not None and self.metric_type in (MetricType.L2, "L2"):
score = 1.0 / (1.0 + raw_distance)
else:
score = raw_distance
memory_obj = OutputData(id=uid, score=score, payload=metadata)
memory.append(memory_obj)
+1 -1
View File
@@ -347,7 +347,7 @@ class PGVector(VectorStoreBase):
)
results = cur.fetchall()
return [OutputData(id=str(r[0]), score=float(r[1]), payload=r[2]) for r in results]
return [OutputData(id=str(r[0]), score=max(0.0, 1.0 - float(r[1])), payload=r[2]) for r in results]
def keyword_search(self, query, top_k=5, filters=None):
"""
+1 -1
View File
@@ -158,7 +158,7 @@ class RedisDB(VectorStoreBase):
return [
MemoryResult(
id=result["memory_id"],
score=float(result["vector_distance"]),
score=max(0.0, 1.0 - float(result["vector_distance"])),
payload={
"hash": result["hash"],
"data": result["memory"],
+3 -1
View File
@@ -80,7 +80,9 @@ class S3Vectors(VectorStoreBase):
except json.JSONDecodeError:
logger.warning(f"Failed to parse metadata for key {v.get('key')}")
payload = {}
results.append(OutputData(id=v.get("key"), score=v.get("distance"), payload=payload))
raw_distance = v.get("distance")
score = max(0.0, 1.0 - raw_distance) if raw_distance is not None else None
results.append(OutputData(id=v.get("key"), score=score, payload=payload))
return results
def insert(self, vectors, payloads=None, ids=None):
+1 -1
View File
@@ -135,7 +135,7 @@ class Supabase(VectorStoreBase):
data=vectors, limit=top_k, filters=filters, include_metadata=True, include_value=True
)
return [OutputData(id=str(result[0]), score=float(result[1]), payload=result[2]) for result in results]
return [OutputData(id=str(result[0]), score=max(0.0, 1.0 - float(result[1])), payload=result[2]) for result in results]
def delete(self, vector_id: str):
"""
+2 -2
View File
@@ -388,8 +388,8 @@ class ValkeyDB(VectorStoreBase):
"""
memory_results = []
for doc in results.docs:
# Extract the score
score = float(doc.vector_score) if hasattr(doc, "vector_score") else None
raw_distance = float(doc.vector_score) if hasattr(doc, "vector_score") else None
score = max(0.0, 1.0 - raw_distance) if raw_distance is not None else None
# Create the payload
payload = {
@@ -114,10 +114,12 @@ class GoogleMatchingEngine(VectorStoreBase):
results = data.get("nearestNeighbors", {}).get("neighbors", [])
output_data = []
for result in results:
raw_distance = result.get("distance")
score = max(0.0, 1.0 - raw_distance) if raw_distance is not None else None
output_data.append(
OutputData(
id=result.get("datapoint").get("datapointId"),
score=result.get("distance"),
score=score,
payload=result.get("datapoint").get("metadata"),
)
)
@@ -264,7 +266,8 @@ class GoogleMatchingEngine(VectorStoreBase):
logger.debug("Adding %s: %s", restrict.name, restrict.allow_tokens[0])
payload[restrict.name] = restrict.allow_tokens[0]
output_data = OutputData(id=neighbor.id, score=neighbor.distance, payload=payload)
score = max(0.0, 1.0 - neighbor.distance) if neighbor.distance is not None else None
output_data = OutputData(id=neighbor.id, score=score, payload=payload)
results.append(output_data)
logger.debug("Returning %d results", len(results))
@@ -413,7 +416,8 @@ class GoogleMatchingEngine(VectorStoreBase):
if restrict.allow_list:
payload[restrict.namespace] = restrict.allow_list[0]
return OutputData(id=neighbor.datapoint.datapoint_id, score=neighbor.distance, payload=payload)
score = max(0.0, 1.0 - neighbor.distance) if neighbor.distance is not None else None
return OutputData(id=neighbor.datapoint.datapoint_id, score=score, payload=payload)
logger.debug("No results found")
return None
+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"}