fix(langchain): search() crashes with TypeError when score is None (#5072)
Co-authored-by: Kartik <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -107,7 +107,7 @@ def score_and_rank(
|
||||
if mem_id is None:
|
||||
continue
|
||||
|
||||
semantic_score = result.get("score", 0.0)
|
||||
semantic_score = result.get("score") or 0.0
|
||||
if semantic_score < threshold:
|
||||
continue
|
||||
|
||||
|
||||
@@ -21,6 +21,16 @@ class OutputData(BaseModel):
|
||||
payload: Optional[Dict] # metadata
|
||||
|
||||
|
||||
# Methods that accept a pre-computed embedding and return (Document, float) pairs.
|
||||
# Tried in order; first match wins. Not part of the base VectorStore contract,
|
||||
# but exposed by several concrete implementations.
|
||||
_SCORED_BY_VECTOR_METHODS = [
|
||||
"similarity_search_by_vector_with_relevance_scores", # Chroma
|
||||
"similarity_search_with_score_by_vector", # FAISS, Qdrant
|
||||
"similarity_search_by_vector_with_score", # Pinecone, YDB
|
||||
]
|
||||
|
||||
|
||||
class Langchain(VectorStoreBase):
|
||||
def __init__(self, client: VectorStore, collection_name: str = "mem0"):
|
||||
self.client = client
|
||||
@@ -95,14 +105,40 @@ class Langchain(VectorStoreBase):
|
||||
"""
|
||||
Search for similar vectors in LangChain.
|
||||
"""
|
||||
# For each vector, perform a similarity search
|
||||
kwargs = {"embedding": vectors, "k": top_k}
|
||||
if filters:
|
||||
results = self.client.similarity_search_by_vector(embedding=vectors, k=top_k, filter=filters)
|
||||
else:
|
||||
results = self.client.similarity_search_by_vector(embedding=vectors, k=top_k)
|
||||
kwargs["filter"] = filters
|
||||
|
||||
final_results = self._parse_output(results)
|
||||
return final_results
|
||||
# Try methods that return (Document, float) pairs — not in the base contract
|
||||
# but available on several concrete implementations.
|
||||
for method_name in _SCORED_BY_VECTOR_METHODS:
|
||||
method = getattr(self.client, method_name, None)
|
||||
if method is None:
|
||||
continue
|
||||
try:
|
||||
results = method(**kwargs)
|
||||
return [
|
||||
OutputData(
|
||||
id=getattr(doc, "id", None),
|
||||
score=float(score),
|
||||
payload=getattr(doc, "metadata", {}),
|
||||
)
|
||||
for doc, score in results
|
||||
]
|
||||
except (NotImplementedError, TypeError):
|
||||
continue
|
||||
|
||||
# Fallback: similarity_search_by_vector returns List[Document] with no scores.
|
||||
# Assign 1.0 so score_and_rank never receives None (None < threshold crashes).
|
||||
docs = self.client.similarity_search_by_vector(**kwargs)
|
||||
return [
|
||||
OutputData(
|
||||
id=getattr(doc, "id", None),
|
||||
score=1.0,
|
||||
payload=getattr(doc, "metadata", {}),
|
||||
)
|
||||
for doc in docs
|
||||
]
|
||||
|
||||
def delete(self, vector_id):
|
||||
"""
|
||||
|
||||
@@ -119,6 +119,13 @@ class TestScoreAndRank:
|
||||
scored = score_and_rank([], {}, {}, threshold=0.1, top_k=10)
|
||||
assert scored == []
|
||||
|
||||
def test_none_score_treated_as_zero(self):
|
||||
"""Defensive: score=None must not crash on None < threshold comparison."""
|
||||
results = [{"id": "a", "score": None, "payload": {"data": "mem a"}}]
|
||||
# Should not raise TypeError; None score is treated as 0.0 and filtered out
|
||||
scored = score_and_rank(results, {}, {}, threshold=0.1, top_k=10)
|
||||
assert scored == []
|
||||
|
||||
def test_score_clamped_to_1(self):
|
||||
results = [{"id": "a", "score": 1.0, "payload": {}}]
|
||||
bm25 = {"a": 1.0}
|
||||
|
||||
@@ -57,6 +57,9 @@ def test_search_vectors(langchain_instance):
|
||||
assert results[0].payload == {"name": "vector1"}
|
||||
assert results[1].id == "id2"
|
||||
assert results[1].payload == {"name": "vector2"}
|
||||
# scores must never be None — score_and_rank crashes on None < threshold
|
||||
assert results[0].score == 1.0
|
||||
assert results[1].score == 1.0
|
||||
|
||||
# Test search with filters
|
||||
filters = {"name": "vector1"}
|
||||
@@ -226,6 +229,56 @@ def test_list_with_exception(langchain_instance):
|
||||
assert results == []
|
||||
|
||||
|
||||
def test_search_score_is_never_none(langchain_instance):
|
||||
"""Regression: similarity_search_by_vector returns Documents with no scores.
|
||||
search() must not propagate None — score_and_rank crashes on None < threshold."""
|
||||
from mem0.utils.scoring import score_and_rank
|
||||
|
||||
mock_docs = [Mock(metadata={"data": "mem A"}, id="id1"), Mock(metadata={"data": "mem B"}, id="id2")]
|
||||
langchain_instance.client.similarity_search_by_vector.return_value = mock_docs
|
||||
|
||||
results = langchain_instance.search(query="test", vectors=[[0.1, 0.2]], top_k=5)
|
||||
|
||||
assert all(r.score is not None for r in results), "score must never be None"
|
||||
# Fallback path: no scored method available, so 1.0 is assigned.
|
||||
assert all(r.score == 1.0 for r in results)
|
||||
|
||||
# Verify the full pipeline does not raise TypeError
|
||||
candidates = [{"id": r.id, "score": r.score, "payload": r.payload} for r in results]
|
||||
ranked = score_and_rank(candidates, {}, {}, threshold=0.1, top_k=5)
|
||||
assert len(ranked) == 2
|
||||
|
||||
|
||||
def test_search_uses_scored_method_when_available(langchain_instance):
|
||||
"""When a scored-by-vector method exists on the client, use it to get real scores."""
|
||||
mock_docs = [Mock(metadata={"data": "mem A"}, id="id1"), Mock(metadata={"data": "mem B"}, id="id2")]
|
||||
# Inject a non-spec method that returns (Document, float) pairs
|
||||
langchain_instance.client.similarity_search_by_vector_with_relevance_scores = Mock(
|
||||
return_value=[(mock_docs[0], 0.95), (mock_docs[1], 0.42)]
|
||||
)
|
||||
|
||||
results = langchain_instance.search(query="test", vectors=[[0.1, 0.2]], top_k=5)
|
||||
|
||||
langchain_instance.client.similarity_search_by_vector_with_relevance_scores.assert_called_once()
|
||||
langchain_instance.client.similarity_search_by_vector.assert_not_called()
|
||||
assert results[0].score == pytest.approx(0.95)
|
||||
assert results[1].score == pytest.approx(0.42)
|
||||
|
||||
|
||||
def test_search_falls_back_when_scored_method_raises_not_implemented(langchain_instance):
|
||||
"""If the scored method raises NotImplementedError, fall back to score=1.0."""
|
||||
mock_docs = [Mock(metadata={"data": "mem A"}, id="id1")]
|
||||
langchain_instance.client.similarity_search_by_vector_with_relevance_scores = Mock(
|
||||
side_effect=NotImplementedError
|
||||
)
|
||||
langchain_instance.client.similarity_search_by_vector.return_value = mock_docs
|
||||
|
||||
results = langchain_instance.search(query="test", vectors=[[0.1, 0.2]], top_k=5)
|
||||
|
||||
langchain_instance.client.similarity_search_by_vector.assert_called_once()
|
||||
assert results[0].score == 1.0
|
||||
|
||||
|
||||
def test_update_wraps_vector_and_payload_in_lists(langchain_instance):
|
||||
"""Regression test for Langchain update() type mismatch.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user