diff --git a/integrations/qdrant/src/haystack_integrations/components/retrievers/qdrant/retriever.py b/integrations/qdrant/src/haystack_integrations/components/retrievers/qdrant/retriever.py index e948249944..9c02dc71bf 100644 --- a/integrations/qdrant/src/haystack_integrations/components/retrievers/qdrant/retriever.py +++ b/integrations/qdrant/src/haystack_integrations/components/retrievers/qdrant/retriever.py @@ -188,9 +188,9 @@ def run( query_embedding=query_embedding, filters=filters, top_k=top_k or self._top_k, - scale_score=scale_score or self._scale_score, - return_embedding=return_embedding or self._return_embedding, - score_threshold=score_threshold or self._score_threshold, + scale_score=self._scale_score if scale_score is None else scale_score, + return_embedding=self._return_embedding if return_embedding is None else return_embedding, + score_threshold=self._score_threshold if score_threshold is None else score_threshold, group_by=group_by or self._group_by, group_size=group_size or self._group_size, ) @@ -243,9 +243,9 @@ async def run_async( query_embedding=query_embedding, filters=filters, top_k=top_k or self._top_k, - scale_score=scale_score or self._scale_score, - return_embedding=return_embedding or self._return_embedding, - score_threshold=score_threshold or self._score_threshold, + scale_score=self._scale_score if scale_score is None else scale_score, + return_embedding=self._return_embedding if return_embedding is None else return_embedding, + score_threshold=self._score_threshold if score_threshold is None else score_threshold, group_by=group_by or self._group_by, group_size=group_size or self._group_size, ) @@ -433,9 +433,9 @@ def run( query_sparse_embedding=query_sparse_embedding, filters=filters, top_k=top_k or self._top_k, - scale_score=scale_score or self._scale_score, - return_embedding=return_embedding or self._return_embedding, - score_threshold=score_threshold or self._score_threshold, + scale_score=self._scale_score if scale_score is None else scale_score, + return_embedding=self._return_embedding if return_embedding is None else return_embedding, + score_threshold=self._score_threshold if score_threshold is None else score_threshold, group_by=group_by or self._group_by, group_size=group_size or self._group_size, ) @@ -493,9 +493,9 @@ async def run_async( query_sparse_embedding=query_sparse_embedding, filters=filters, top_k=top_k or self._top_k, - scale_score=scale_score or self._scale_score, - return_embedding=return_embedding or self._return_embedding, - score_threshold=score_threshold or self._score_threshold, + scale_score=self._scale_score if scale_score is None else scale_score, + return_embedding=self._return_embedding if return_embedding is None else return_embedding, + score_threshold=self._score_threshold if score_threshold is None else score_threshold, group_by=group_by or self._group_by, group_size=group_size or self._group_size, ) @@ -704,8 +704,8 @@ def run( query_sparse_embedding=query_sparse_embedding, filters=filters, top_k=top_k or self._top_k, - return_embedding=return_embedding or self._return_embedding, - score_threshold=score_threshold or self._score_threshold, + return_embedding=self._return_embedding if return_embedding is None else return_embedding, + score_threshold=self._score_threshold if score_threshold is None else score_threshold, group_by=group_by or self._group_by, group_size=group_size or self._group_size, rrf_k=rrf_k if rrf_k is not None else self._rrf_k, @@ -774,8 +774,8 @@ async def run_async( query_sparse_embedding=query_sparse_embedding, filters=filters, top_k=top_k or self._top_k, - return_embedding=return_embedding or self._return_embedding, - score_threshold=score_threshold or self._score_threshold, + return_embedding=self._return_embedding if return_embedding is None else return_embedding, + score_threshold=self._score_threshold if score_threshold is None else score_threshold, group_by=group_by or self._group_by, group_size=group_size or self._group_size, rrf_k=rrf_k if rrf_k is not None else self._rrf_k, diff --git a/integrations/qdrant/tests/test_embedding_retriever.py b/integrations/qdrant/tests/test_embedding_retriever.py index a3f9e52577..a44bb2b2ff 100644 --- a/integrations/qdrant/tests/test_embedding_retriever.py +++ b/integrations/qdrant/tests/test_embedding_retriever.py @@ -200,6 +200,37 @@ async def test_run_async(self): mock_store._query_by_embedding_async.assert_awaited_once() assert res["documents"][0].content == "doc" + def test_run_falsy_runtime_values_override_init(self): + mock_store = Mock(spec=QdrantDocumentStore) + mock_store._query_by_embedding.return_value = [] + + retriever = QdrantEmbeddingRetriever( + document_store=mock_store, scale_score=True, return_embedding=True, score_threshold=0.5 + ) + retriever.run(query_embedding=[0.5, 0.7], scale_score=False, return_embedding=False, score_threshold=0.0) + + call_kwargs = mock_store._query_by_embedding.call_args.kwargs + assert call_kwargs["scale_score"] is False + assert call_kwargs["return_embedding"] is False + assert call_kwargs["score_threshold"] == 0.0 + + @pytest.mark.asyncio + async def test_run_async_falsy_runtime_values_override_init(self): + mock_store = Mock(spec=QdrantDocumentStore) + mock_store._query_by_embedding_async = AsyncMock(return_value=[]) + + retriever = QdrantEmbeddingRetriever( + document_store=mock_store, scale_score=True, return_embedding=True, score_threshold=0.5 + ) + await retriever.run_async( + query_embedding=[0.5, 0.7], scale_score=False, return_embedding=False, score_threshold=0.0 + ) + + call_kwargs = mock_store._query_by_embedding_async.call_args.kwargs + assert call_kwargs["scale_score"] is False + assert call_kwargs["return_embedding"] is False + assert call_kwargs["score_threshold"] == 0.0 + def test_run_raises_when_merge_with_native_init_filter(self): document_store = QdrantDocumentStore(location=":memory:", index="test") retriever = QdrantEmbeddingRetriever( diff --git a/integrations/qdrant/tests/test_hybrid_retriever.py b/integrations/qdrant/tests/test_hybrid_retriever.py index 7e993a836b..b8f5139d01 100644 --- a/integrations/qdrant/tests/test_hybrid_retriever.py +++ b/integrations/qdrant/tests/test_hybrid_retriever.py @@ -229,6 +229,39 @@ def test_run_runtime_rrf_params_override_init(self): assert call_args[1]["rrf_k"] == 100 assert call_args[1]["rrf_weights"] == [3.0, 1.0] + def test_run_falsy_runtime_values_override_init(self): + mock_store = Mock(spec=QdrantDocumentStore) + mock_store._query_hybrid.return_value = [] + + retriever = QdrantHybridRetriever(document_store=mock_store, return_embedding=True, score_threshold=0.5) + retriever.run( + query_embedding=[0.5, 0.7], + query_sparse_embedding=SparseEmbedding(indices=[0, 5], values=[0.1, 0.7]), + return_embedding=False, + score_threshold=0.0, + ) + + call_args = mock_store._query_hybrid.call_args + assert call_args[1]["return_embedding"] is False + assert call_args[1]["score_threshold"] == 0.0 + + @pytest.mark.asyncio + async def test_run_async_falsy_runtime_values_override_init(self): + mock_store = Mock(spec=QdrantDocumentStore) + mock_store._query_hybrid_async = AsyncMock(return_value=[]) + + retriever = QdrantHybridRetriever(document_store=mock_store, return_embedding=True, score_threshold=0.5) + await retriever.run_async( + query_embedding=[0.5, 0.7], + query_sparse_embedding=SparseEmbedding(indices=[0, 5], values=[0.1, 0.7]), + return_embedding=False, + score_threshold=0.0, + ) + + call_args = mock_store._query_hybrid_async.call_args + assert call_args[1]["return_embedding"] is False + assert call_args[1]["score_threshold"] == 0.0 + def test_run_with_group_by(self): mock_store = Mock(spec=QdrantDocumentStore) sparse_embedding = SparseEmbedding(indices=[0, 1, 2, 3], values=[0.1, 0.8, 0.05, 0.33]) diff --git a/integrations/qdrant/tests/test_sparse_embedding_retriever.py b/integrations/qdrant/tests/test_sparse_embedding_retriever.py index 4243b0a5a0..12c36e3486 100644 --- a/integrations/qdrant/tests/test_sparse_embedding_retriever.py +++ b/integrations/qdrant/tests/test_sparse_embedding_retriever.py @@ -199,6 +199,39 @@ async def test_run_async(self): mock_store._query_by_sparse_async.assert_awaited_once() assert res["documents"][0].content == "doc" + def test_run_falsy_runtime_values_override_init(self): + mock_store = Mock(spec=QdrantDocumentStore) + mock_store._query_by_sparse.return_value = [] + sparse = SparseEmbedding(indices=[0, 5], values=[0.1, 0.7]) + + retriever = QdrantSparseEmbeddingRetriever( + document_store=mock_store, scale_score=True, return_embedding=True, score_threshold=0.5 + ) + retriever.run(query_sparse_embedding=sparse, scale_score=False, return_embedding=False, score_threshold=0.0) + + call_kwargs = mock_store._query_by_sparse.call_args.kwargs + assert call_kwargs["scale_score"] is False + assert call_kwargs["return_embedding"] is False + assert call_kwargs["score_threshold"] == 0.0 + + @pytest.mark.asyncio + async def test_run_async_falsy_runtime_values_override_init(self): + mock_store = Mock(spec=QdrantDocumentStore) + mock_store._query_by_sparse_async = AsyncMock(return_value=[]) + sparse = SparseEmbedding(indices=[0, 5], values=[0.1, 0.7]) + + retriever = QdrantSparseEmbeddingRetriever( + document_store=mock_store, scale_score=True, return_embedding=True, score_threshold=0.5 + ) + await retriever.run_async( + query_sparse_embedding=sparse, scale_score=False, return_embedding=False, score_threshold=0.0 + ) + + call_kwargs = mock_store._query_by_sparse_async.call_args.kwargs + assert call_kwargs["scale_score"] is False + assert call_kwargs["return_embedding"] is False + assert call_kwargs["score_threshold"] == 0.0 + def test_run_raises_when_merge_with_native_filter(self): document_store = QdrantDocumentStore(location=":memory:", index="test") retriever = QdrantSparseEmbeddingRetriever(