From b0b800fbd8a8415e64438a663cbf3ad3d748e66f Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 6 Feb 2026 15:52:09 +0800 Subject: [PATCH] feat(vector-store): add threshold-based filtering to vector search --- reme/core/vector_store/base_vector_store.py | 17 +- reme/core/vector_store/chroma_vector_store.py | 26 ++- reme/core/vector_store/es_vector_store.py | 34 +++- reme/core/vector_store/local_vector_store.py | 23 ++- reme/core/vector_store/pgvector_store.py | 31 +++- reme/core/vector_store/qdrant_vector_store.py | 25 ++- tests/test_vector_store.py | 175 ++++++++++++++++++ 7 files changed, 304 insertions(+), 27 deletions(-) diff --git a/reme/core/vector_store/base_vector_store.py b/reme/core/vector_store/base_vector_store.py index 62a73a3e..addaa2ea 100644 --- a/reme/core/vector_store/base_vector_store.py +++ b/reme/core/vector_store/base_vector_store.py @@ -73,8 +73,21 @@ class BaseVectorStore(ABC): """Add one or more vector nodes into the current collection.""" @abstractmethod - async def search(self, query: str, limit: int = 5, filters: dict | None = None, **kwargs) -> list[VectorNode]: - """Find the most similar vector nodes based on a text query.""" + async def search( + self, + query: str, + limit: int = 5, + candidates: int | None = None, + filters: dict | None = None, + threshold: float | None = None, + **kwargs, + ) -> list[VectorNode]: + """Find the most similar vector nodes based on a text query. + + When threshold is None, uses default behavior. + When threshold is set, searches max(candidates, limit) nodes, + filters by threshold, then returns top limit results. + """ @abstractmethod async def delete(self, vector_ids: str | list[str], **kwargs) -> None: diff --git a/reme/core/vector_store/chroma_vector_store.py b/reme/core/vector_store/chroma_vector_store.py index 5b0748f5..76b4b81a 100644 --- a/reme/core/vector_store/chroma_vector_store.py +++ b/reme/core/vector_store/chroma_vector_store.py @@ -307,21 +307,35 @@ class ChromaVectorStore(BaseVectorStore): self, query: str, limit: int = 5, + candidates: int | None = None, filters: dict | None = None, + threshold: float | None = None, **kwargs, ) -> list[VectorNode]: - """Search for the most similar vector nodes based on a text query.""" + """Search for the most similar vector nodes based on a text query. + + When threshold is None, uses default behavior. + When threshold is set, searches max(candidates, limit) nodes, + filters by threshold, then returns top limit results. + """ query_vector = await self.get_embedding(query) where_clause = self._generate_where_clause(filters) include_embeddings = kwargs.get("include_embeddings", False) + # When threshold is set, search more candidates + if threshold is not None: + effective_candidates = candidates if candidates is not None else limit * 2 + search_limit = max(effective_candidates, limit) + else: + search_limit = limit + def _search(): include: list = ["documents", "metadatas", "distances"] if include_embeddings: include.append("embeddings") return self.collection.query( query_embeddings=[query_vector], - n_results=limit, + n_results=search_limit, where=where_clause, include=include, ) @@ -329,9 +343,11 @@ class ChromaVectorStore(BaseVectorStore): results = await self._run_sync_in_executor(_search) nodes = self._parse_results(results, include_score=True) - score_threshold = kwargs.get("score_threshold") - if score_threshold is not None: - nodes = [n for n in nodes if n.metadata.get("score", 0) >= score_threshold] + # Apply threshold filtering if specified + if threshold is not None: + nodes = [n for n in nodes if n.metadata.get("score", 0) >= threshold] + nodes = nodes[:limit] + return nodes async def delete(self, vector_ids: str | list[str], **kwargs): diff --git a/reme/core/vector_store/es_vector_store.py b/reme/core/vector_store/es_vector_store.py index 62b28bc1..f9cfad7d 100644 --- a/reme/core/vector_store/es_vector_store.py +++ b/reme/core/vector_store/es_vector_store.py @@ -238,31 +238,48 @@ class ESVectorStore(BaseVectorStore): self, query: str, limit: int = 5, + candidates: int | None = None, filters: dict | None = None, + threshold: float | None = None, **kwargs, ) -> list[VectorNode]: """Perform a kNN similarity search based on a text query. + When threshold is None, uses default behavior. + When threshold is set, searches max(candidates, limit) nodes, + filters by threshold, then returns top limit results. + Args: query: The text to search for. limit: Maximum number of nearest neighbors to return. - filters: Metadata filters for exact match or 'IN' operations. - **kwargs: Search parameters like num_candidates or score_threshold. + candidates: Number of candidates to consider during search. + filters: Metadata filters for exact match or range operations. + threshold: Minimum score threshold for results. + **kwargs: Additional search parameters. Returns: List of VectorNode objects ordered by similarity. """ query_vector = await self.get_embedding(query) - num_candidates = kwargs.get("num_candidates", limit * 2) + + # When threshold is set, search more candidates + if threshold is not None: + effective_candidates = candidates if candidates is not None else limit * 2 + search_candidates = max(effective_candidates, limit) + search_limit = search_candidates + else: + effective_candidates = candidates if candidates is not None else limit * 2 + search_candidates = effective_candidates + search_limit = limit search_query: dict = { "knn": { "field": "vector", "query_vector": query_vector, - "k": limit, - "num_candidates": num_candidates, + "k": search_limit, + "num_candidates": search_candidates, }, - "size": limit, + "size": search_limit, } if filters: @@ -300,6 +317,11 @@ class ESVectorStore(BaseVectorStore): node.metadata["score"] = hit["_score"] results.append(node) + # Apply threshold filtering if specified + if threshold is not None: + results = [n for n in results if n.metadata.get("score", 0) >= threshold] + results = results[:limit] + return results async def delete(self, vector_ids: str | list[str], refresh: bool = True, **kwargs): diff --git a/reme/core/vector_store/local_vector_store.py b/reme/core/vector_store/local_vector_store.py index c117ddf0..22db669e 100644 --- a/reme/core/vector_store/local_vector_store.py +++ b/reme/core/vector_store/local_vector_store.py @@ -193,10 +193,17 @@ class LocalVectorStore(BaseVectorStore): self, query: str, limit: int = 5, + candidates: int | None = None, filters: dict | None = None, + threshold: float | None = None, **kwargs, ) -> list[VectorNode]: - """Search for nodes similar to the query using brute-force cosine similarity.""" + """Search for nodes similar to the query using brute-force cosine similarity. + + When threshold is None, uses default behavior. + When threshold is set, searches max(candidates, limit) nodes, + filters by threshold, then returns top limit results. + """ query_vector = await self.get_embedding(query) all_nodes = self._load_all_nodes() filtered_nodes = [node for node in all_nodes if self._match_filters(node, filters)] @@ -215,11 +222,17 @@ class LocalVectorStore(BaseVectorStore): scored_nodes.sort(key=lambda x: x[1], reverse=True) - score_threshold = kwargs.get("score_threshold") - if score_threshold is not None: - scored_nodes = [(node, score) for node, score in scored_nodes if score >= score_threshold] + # Apply threshold filtering if specified + if threshold is not None: + # When threshold is set, consider more candidates + effective_candidates = candidates if candidates is not None else limit * 2 + search_limit = max(effective_candidates, limit) + scored_nodes = scored_nodes[:search_limit] + scored_nodes = [(node, score) for node, score in scored_nodes if score >= threshold] + scored_nodes = scored_nodes[:limit] + else: + scored_nodes = scored_nodes[:limit] - scored_nodes = scored_nodes[:limit] results = [] for node, score in scored_nodes: node.metadata["score"] = score diff --git a/reme/core/vector_store/pgvector_store.py b/reme/core/vector_store/pgvector_store.py index 0ec7d6a8..dcb49201 100644 --- a/reme/core/vector_store/pgvector_store.py +++ b/reme/core/vector_store/pgvector_store.py @@ -333,10 +333,17 @@ class PGVectorStore(BaseVectorStore): self, query: str, limit: int = 5, + candidates: int | None = None, filters: dict | None = None, + threshold: float | None = None, **kwargs, ) -> list[VectorNode]: - """Perform vector similarity search with optional metadata filtering.""" + """Perform vector similarity search with optional metadata filtering. + + When threshold is None, uses default behavior. + When threshold is set, searches max(candidates, limit) nodes, + filters by threshold, then returns top limit results. + """ await self._ensure_collection_exists() query_vector = await self.get_embedding(query) @@ -351,6 +358,13 @@ class PGVectorStore(BaseVectorStore): new_placeholder = f"${i + 1}" filter_clause = re.sub(rf"\${i}\b", new_placeholder, filter_clause) + # When threshold is set, search more candidates + if threshold is not None: + effective_candidates = candidates if candidates is not None else limit * 2 + search_limit = max(effective_candidates, limit) + else: + search_limit = limit + async with pool.acquire() as conn: sql = f""" SELECT id, content, vector, metadata, vector <=> $1::vector AS distance @@ -359,14 +373,17 @@ class PGVectorStore(BaseVectorStore): ORDER BY distance LIMIT ${len(filter_params) + 2} """ - rows = await conn.fetch(sql, vector_str, *filter_params, limit) + rows = await conn.fetch(sql, vector_str, *filter_params, search_limit) results = [] - score_threshold = kwargs.get("score_threshold") for row in rows: distance = row["distance"] - if score_threshold is not None and distance > score_threshold: + # Convert distance to score (1 - distance for cosine) + score = 1 - distance + + # Apply threshold filtering if specified + if threshold is not None and score < threshold: continue vector_data = None @@ -379,7 +396,7 @@ class PGVectorStore(BaseVectorStore): if isinstance(metadata, str): metadata = json.loads(metadata) - metadata["score"] = 1 - distance + metadata["score"] = score metadata["_distance"] = distance node = VectorNode( @@ -390,6 +407,10 @@ class PGVectorStore(BaseVectorStore): ) results.append(node) + # Apply limit after threshold filtering + if threshold is not None: + results = results[:limit] + return results async def delete(self, vector_ids: str | list[str], **kwargs): diff --git a/reme/core/vector_store/qdrant_vector_store.py b/reme/core/vector_store/qdrant_vector_store.py index 93ccee70..f14dd958 100644 --- a/reme/core/vector_store/qdrant_vector_store.py +++ b/reme/core/vector_store/qdrant_vector_store.py @@ -325,20 +325,33 @@ class QdrantVectorStore(BaseVectorStore): self, query: str, limit: int = 5, + candidates: int | None = None, filters: dict | None = None, + threshold: float | None = None, **kwargs: Any, ) -> list[VectorNode]: - """Search for the most similar vectors based on a text query.""" + """Search for the most similar vectors based on a text query. + + When threshold is None, uses default behavior. + When threshold is set, searches max(candidates, limit) nodes, + filters by threshold, then returns top limit results. + """ query_vector = await self.get_embedding(query) query_filter = self._create_filter(filters) if filters else None - score_threshold = kwargs.get("score_threshold", None) + + # When threshold is set, search more candidates + if threshold is not None: + effective_candidates = candidates if candidates is not None else limit * 2 + search_limit = max(effective_candidates, limit) + else: + search_limit = limit results = await self.client.query_points( collection_name=self.collection_name, query=query_vector, query_filter=query_filter, - limit=limit, - score_threshold=score_threshold, + limit=search_limit, + score_threshold=threshold, ) nodes = [] @@ -353,6 +366,10 @@ class QdrantVectorStore(BaseVectorStore): node.metadata["score"] = point.score nodes.append(node) + # Apply limit after threshold filtering (Qdrant already filters by threshold) + if threshold is not None: + nodes = nodes[:limit] + return nodes async def delete(self, vector_ids: str | list[str], **kwargs: Any): diff --git a/tests/test_vector_store.py b/tests/test_vector_store.py index 14ed1c1d..d004c43f 100644 --- a/tests/test_vector_store.py +++ b/tests/test_vector_store.py @@ -414,6 +414,180 @@ async def test_search_with_multiple_filters(store: BaseVectorStore, _store_name: logger.info("✓ Multiple filters search test passed") +async def test_search_with_threshold(store: BaseVectorStore, _store_name: str): + """Test vector search with threshold parameter. + + When threshold is None, uses default behavior. + When threshold is set, searches max(candidates, limit) nodes, + filters by threshold, then returns top limit results. + """ + logger.info("=" * 20 + " SEARCH WITH THRESHOLD TEST " + "=" * 20) + + # Clean up any existing test data first + try: + existing_nodes = await store.list(filters={"test_type": "threshold_test"}) + if existing_nodes: + await store.delete([node.vector_id for node in existing_nodes]) + logger.info(f"Cleaned up {len(existing_nodes)} existing test nodes") + except Exception as e: + logger.warning(f"Failed to clean up existing nodes: {e}") + + # Insert test nodes with varied content for different similarity scores + threshold_nodes = [ + VectorNode( + vector_id="threshold_node_1", + content="Machine learning and artificial intelligence are the future of technology.", + metadata={"test_type": "threshold_test", "relevance": "high"}, + ), + VectorNode( + vector_id="threshold_node_2", + content="Deep learning uses neural networks for artificial intelligence applications.", + metadata={"test_type": "threshold_test", "relevance": "high"}, + ), + VectorNode( + vector_id="threshold_node_3", + content="Natural language processing is a branch of AI technology.", + metadata={"test_type": "threshold_test", "relevance": "medium"}, + ), + VectorNode( + vector_id="threshold_node_4", + content="Computer vision enables machines to interpret visual information.", + metadata={"test_type": "threshold_test", "relevance": "medium"}, + ), + VectorNode( + vector_id="threshold_node_5", + content="Software development requires good coding practices.", + metadata={"test_type": "threshold_test", "relevance": "low"}, + ), + VectorNode( + vector_id="threshold_node_6", + content="Database management systems store and retrieve data efficiently.", + metadata={"test_type": "threshold_test", "relevance": "low"}, + ), + VectorNode( + vector_id="threshold_node_7", + content="Cloud computing provides scalable infrastructure for applications.", + metadata={"test_type": "threshold_test", "relevance": "low"}, + ), + VectorNode( + vector_id="threshold_node_8", + content="Cooking pasta requires boiling water and adding salt.", + metadata={"test_type": "threshold_test", "relevance": "none"}, + ), + ] + + await store.insert(threshold_nodes) + logger.info(f"✓ Inserted {len(threshold_nodes)} test nodes") + + # Test 1: Search without threshold (baseline) + results_baseline = await store.search( + query="artificial intelligence and machine learning", + limit=3, + candidates=10, + filters={"test_type": "threshold_test"}, + ) + logger.info(f"Test 1 - Search without threshold: {len(results_baseline)} results") + for i, r in enumerate(results_baseline, 1): + score = r.metadata.get("score", 0) + logger.info(f" Result {i}: {r.vector_id}, score={score:.4f}") + + assert len(results_baseline) <= 3, "Should return at most limit results" + logger.info("✓ Baseline search validated") + + # Test 2: Search with threshold (should filter low-score results) + threshold_value = 0.5 + results_with_threshold = await store.search( + query="artificial intelligence and machine learning", + limit=3, + candidates=10, + filters={"test_type": "threshold_test"}, + threshold=threshold_value, + ) + logger.info( + f"Test 2 - Search with threshold={threshold_value}: {len(results_with_threshold)} results", + ) + for i, r in enumerate(results_with_threshold, 1): + score = r.metadata.get("score", 0) + logger.info(f" Result {i}: {r.vector_id}, score={score:.4f}") + # Verify all results meet threshold + assert score >= threshold_value, f"Score {score} should be >= threshold {threshold_value}" + + # Should return at most limit results + assert len(results_with_threshold) <= 3, "Should return at most limit results" + logger.info("✓ Threshold filtering validated") + + # Test 3: Search with high threshold (should filter more strictly) + high_threshold = 0.7 + results_high_threshold = await store.search( + query="artificial intelligence and machine learning", + limit=10, # Use same limit as candidates to get all results above threshold + candidates=10, + filters={"test_type": "threshold_test"}, + threshold=high_threshold, + ) + logger.info( + f"Test 3 - Search with high threshold={high_threshold}: {len(results_high_threshold)} results", + ) + for i, r in enumerate(results_high_threshold, 1): + score = r.metadata.get("score", 0) + logger.info(f" Result {i}: {r.vector_id}, score={score:.4f}") + assert score >= high_threshold, f"Score {score} should be >= threshold {high_threshold}" + + # Compare with same limit and lower threshold to verify stricter filtering + results_low_threshold = await store.search( + query="artificial intelligence and machine learning", + limit=10, + candidates=10, + filters={"test_type": "threshold_test"}, + threshold=0.3, # Lower threshold + ) + # High threshold should return same or fewer results than low threshold + assert len(results_high_threshold) <= len( + results_low_threshold, + ), ( + f"High threshold ({len(results_high_threshold)}) should return fewer results than low threshold " + f"({len(results_low_threshold)})" + ) + logger.info("✓ High threshold filtering validated") + + # Test 4: Search with threshold and large candidates + results_large_candidates = await store.search( + query="artificial intelligence and machine learning", + limit=2, + candidates=20, + filters={"test_type": "threshold_test"}, + threshold=0.5, + ) + logger.info( + f"Test 4 - Search with threshold and large candidates: {len(results_large_candidates)} results", + ) + # Should search max(candidates, limit) = 20 nodes, filter by threshold, return top 2 + assert len(results_large_candidates) <= 2, "Should return at most limit results" + for r in results_large_candidates: + score = r.metadata.get("score", 0) + assert score >= 0.5, f"Score {score} should be >= 0.5" + logger.info("✓ Large candidates with threshold validated") + + # Test 5: Verify that threshold=None behaves same as no threshold + results_none_threshold = await store.search( + query="artificial intelligence and machine learning", + limit=3, + candidates=10, + filters={"test_type": "threshold_test"}, + threshold=None, + ) + logger.info(f"Test 5 - Search with threshold=None: {len(results_none_threshold)} results") + # Should behave same as baseline (Test 1) + assert len(results_none_threshold) <= 3, "Should return at most limit results" + logger.info("✓ threshold=None validated") + + # Clean up test data + await store.delete([node.vector_id for node in threshold_nodes]) + logger.info("Cleaned up test nodes") + + logger.info("✓ Search with threshold test passed") + + async def test_get_by_id(store: BaseVectorStore, _store_name: str): """Test retrieving nodes by vector_id (single and batch).""" logger.info("=" * 20 + " GET BY ID TEST " + "=" * 20) @@ -1671,6 +1845,7 @@ async def run_all_tests_for_store(store_type: str, store_name: str): await test_search_with_single_filter(store, store_name) await test_search_with_exact_match_filter(store, store_name) await test_search_with_multiple_filters(store, store_name) + await test_search_with_threshold(store, store_name) await test_get_by_id(store, store_name) await test_list_all(store, store_name) await test_list_with_filters(store, store_name)