feat(vector-store): add threshold-based filtering to vector search

This commit is contained in:
jinli.yl 2026-02-06 15:52:09 +08:00
parent 9e50e71b52
commit b0b800fbd8
7 changed files with 304 additions and 27 deletions

View file

@ -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:

View file

@ -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):

View file

@ -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):

View file

@ -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

View file

@ -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):

View file

@ -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):

View file

@ -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)