diff --git a/mem0/vector_stores/baidu.py b/mem0/vector_stores/baidu.py index e9437232d..2831fcaf7 100644 --- a/mem0/vector_stores/baidu.py +++ b/mem0/vector_stores/baidu.py @@ -224,8 +224,15 @@ class BaiduDB(VectorStoreBase): output = [] for row in res.rows: row_data = row.get("row", {}) + # Mochow returns the raw L2 distance (lower = closer). Convert it to a + # similarity score (higher = better) to satisfy the VectorStoreBase + # contract, mirroring the milvus provider. Non-L2 metrics already + # return a higher-is-better score. + raw_score = row.get("score", 0.0) + if self.metric_type in (MetricType.L2, "L2"): + raw_score = 1.0 / (1.0 + raw_score) output_data = OutputData( - id=row_data.get("id"), score=row.get("score", 0.0), payload=row_data.get("metadata", {}) + id=row_data.get("id"), score=raw_score, payload=row_data.get("metadata", {}) ) output.append(output_data) diff --git a/tests/vector_stores/test_baidu.py b/tests/vector_stores/test_baidu.py index a3777a1d8..8ebf112ec 100644 --- a/tests/vector_stores/test_baidu.py +++ b/tests/vector_stores/test_baidu.py @@ -137,6 +137,27 @@ def test_search(mochow_instance, mock_mochow_client): assert results[1].payload == {"name": "vector2"} +def test_search_converts_l2_distance_to_similarity(mochow_instance): + # On the L2 metric, Mochow returns raw distances (lower = closer). search() + # must convert them to similarity scores (higher = better) to satisfy the + # VectorStoreBase contract, mirroring the milvus provider. + mochow_instance.metric_type = "L2" + + mock_search_results = Mock() + mock_search_results.rows = [ + {"row": {"id": "id1", "metadata": {"name": "vector1"}}, "score": 0.5}, + {"row": {"id": "id2", "metadata": {"name": "vector2"}}, "score": 2.0}, + ] + mochow_instance._table.vector_search.return_value = mock_search_results + + results = mochow_instance.search(query="test", vectors=[0.1, 0.2, 0.3], top_k=2) + + # 1.0 / (1.0 + distance): the closer memory must score higher than the far one. + assert results[0].score == pytest.approx(1.0 / 1.5) + assert results[1].score == pytest.approx(1.0 / 3.0) + assert results[0].score > results[1].score + + def test_search_with_filters(mochow_instance, mock_mochow_client): mochow_instance._table.vector_search.return_value = Mock(rows=[])