fix(vector_stores): normalize scores to similarity (higher = better) across all backends (#5391)
This commit is contained in:
@@ -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')
|
||||
|
||||
Reference in New Issue
Block a user