feat(vector-store): add sorting capability to list operations across all backends

This commit is contained in:
jinli.yl 2026-01-09 19:50:42 +08:00
parent 2d98169d80
commit a172a6854b
9 changed files with 269 additions and 12 deletions

View file

@ -129,7 +129,7 @@ python run_exp_statistic.py
- `best@k`: Takes groups of k runs per task, finds the maximum score in each group, then averages these maximums
- `pass@k`: Takes groups of k runs per task, measures the probability that at least one out of k independent task runs is successful.
- Higher k values show potential performance, lower k values show consistency
- In our AppWorld experiments, we report Task Goal Completion (TGC) metric, which measures percentage of tasks for which the agent passes all evaluation tests.
- In our AppWorld experiments, we report Task Goal Completion (TGC) metric, which measures percentage of tasks for which the agent passes all evaluation tests.
**Output Files**

View file

@ -30,7 +30,7 @@ embedding_model:
vector_store:
default:
backend: local
backend: chroma
embedding_model: default
collection_name: reme

View file

@ -89,8 +89,21 @@ class BaseVectorStore(ABC):
"""Fetch specific vector nodes from the collection by their IDs."""
@abstractmethod
async def list(self, filters: dict | None = None, limit: int | None = None) -> list[VectorNode]:
"""Retrieve vectors from the collection that match the given filters."""
async def list(
self,
filters: dict | None = None,
limit: int | None = None,
sort_key: str | None = None,
reverse: bool = False,
) -> list[VectorNode]:
"""Retrieve vectors from the collection that match the given filters.
Args:
filters: Dictionary of filter conditions to match vectors
limit: Maximum number of vectors to return
sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting
reverse: If True, sort in descending order; if False, sort in ascending order
"""
async def close(self) -> None:
"""Release resources and close active connections to the vector store."""

View file

@ -360,19 +360,49 @@ class ChromaVectorStore(BaseVectorStore):
self,
filters: dict | None = None,
limit: int | None = None,
sort_key: str | None = None,
reverse: bool = False,
) -> list[VectorNode]:
"""List vector nodes matching optional metadata filters."""
"""List vector nodes matching optional metadata filters.
Args:
filters: Dictionary of filter conditions to match vectors
limit: Maximum number of vectors to return
sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting
reverse: If True, sort in descending order; if False, sort in ascending order
"""
where_clause = self._generate_where_clause(filters)
# If sorting is needed, fetch all records first, then apply limit after sorting
fetch_limit = None if sort_key else limit
def _list():
return self.collection.get(
where=where_clause,
limit=limit,
limit=fetch_limit,
include=["documents", "metadatas", "embeddings"],
)
results = await self._run_sync_in_executor(_list)
return self._parse_results(results)
nodes = self._parse_results(results)
# Apply sorting if sort_key is provided
if sort_key:
# Sort with proper handling of None and missing values
def sort_key_func(node):
value = node.metadata.get(sort_key)
if value is None:
# Return appropriate default based on reverse flag
return float("-inf") if not reverse else float("inf")
return value
nodes.sort(key=sort_key_func, reverse=reverse)
# Apply limit after sorting
if limit is not None:
nodes = nodes[:limit]
return nodes
async def count(self) -> int:
"""Return the total number of vectors in the current collection."""

View file

@ -411,12 +411,16 @@ class ESVectorStore(BaseVectorStore):
self,
filters: dict | None = None,
limit: int | None = None,
sort_key: str | None = None,
reverse: bool = False,
) -> list[VectorNode]:
"""Retrieve a list of nodes filtered by metadata or limit.
Args:
filters: Optional metadata filtering criteria.
limit: Maximum number of nodes to return.
sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting
reverse: If True, sort in descending order; if False, sort in ascending order
Returns:
A list of matching VectorNode objects.
@ -432,6 +436,16 @@ class ESVectorStore(BaseVectorStore):
filter_conditions.append({"term": {f"metadata.{key}": value}})
query["query"] = {"bool": {"must": filter_conditions}}
# Add sorting to the Elasticsearch query if sort_key is provided
if sort_key:
query["sort"] = [
{
f"metadata.{sort_key}": {
"order": "desc" if reverse else "asc",
},
},
]
if limit:
query["size"] = limit
else:

View file

@ -266,11 +266,32 @@ class LocalVectorStore(BaseVectorStore):
self,
filters: dict | None = None,
limit: int | None = None,
sort_key: str | None = None,
reverse: bool = False,
) -> list[VectorNode]:
"""List vector nodes in the collection with optional filtering and limits."""
"""List vector nodes in the collection with optional filtering and limits.
Args:
filters: Dictionary of filter conditions to match vectors
limit: Maximum number of vectors to return
sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting
reverse: If True, sort in descending order; if False, sort in ascending order
"""
all_nodes = self._load_all_nodes()
filtered_nodes = [node for node in all_nodes if self._match_filters(node, filters)]
# Apply sorting if sort_key is provided
if sort_key:
# Sort with proper handling of None and missing values
def sort_key_func(node):
value = node.metadata.get(sort_key)
if value is None:
# Return appropriate default based on reverse flag
return float("-inf") if not reverse else float("inf")
return value
filtered_nodes.sort(key=sort_key_func, reverse=reverse)
if limit is not None:
filtered_nodes = filtered_nodes[:limit]

View file

@ -455,13 +455,28 @@ class PGVectorStore(BaseVectorStore):
self,
filters: dict | None = None,
limit: int | None = None,
sort_key: str | None = None,
reverse: bool = False,
) -> list[VectorNode]:
"""Return a list of vector nodes matching the provided filters and limit."""
"""Return a list of vector nodes matching the provided filters and limit.
Args:
filters: Dictionary of filter conditions to match vectors
limit: Maximum number of vectors to return
sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting
reverse: If True, sort in descending order; if False, sort in ascending order
"""
await self._ensure_collection_exists()
pool = await self._get_pool()
filter_clause, filter_params = self._build_filter_clause(filters)
# Add ORDER BY clause if sort_key is provided
order_clause = ""
if sort_key:
order_direction = "DESC" if reverse else "ASC"
order_clause = f"ORDER BY metadata->>'{sort_key}' {order_direction}"
limit_clause = ""
if limit:
limit_clause = f"LIMIT ${len(filter_params) + 1}"
@ -472,6 +487,7 @@ class PGVectorStore(BaseVectorStore):
SELECT id, content, vector, metadata
FROM {self.collection_name}
{filter_clause}
{order_clause}
{limit_clause}
"""
rows = await conn.fetch(sql, *filter_params)

View file

@ -412,15 +412,26 @@ class QdrantVectorStore(BaseVectorStore):
self,
filters: dict | None = None,
limit: int | None = None,
sort_key: str | None = None,
reverse: bool = False,
) -> list[VectorNode]:
"""List all vector nodes in the collection matching the filter criteria."""
"""List all vector nodes in the collection matching the filter criteria.
Args:
filters: Dictionary of filter conditions to match vectors
limit: Maximum number of vectors to return
sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting
reverse: If True, sort in descending order; if False, sort in ascending order
"""
scroll_filter = self._create_filter(filters) if filters else None
limit = limit or 10000
# If sorting is needed, fetch more records than the limit to ensure correct sorting
fetch_limit = 10000 if sort_key else (limit or 10000)
records, _ = await self.client.scroll(
collection_name=self.collection_name,
scroll_filter=scroll_filter,
limit=limit,
limit=fetch_limit,
with_payload=True,
with_vectors=True,
)
@ -436,6 +447,22 @@ class QdrantVectorStore(BaseVectorStore):
)
results.append(node)
# Apply sorting if sort_key is provided
if sort_key:
# Sort with proper handling of None and missing values
def sort_key_func(node):
value = node.metadata.get(sort_key)
if value is None:
# Return appropriate default based on reverse flag
return float("-inf") if not reverse else float("inf")
return value
results.sort(key=sort_key_func, reverse=reverse)
# Apply limit after sorting
if limit is not None:
results = results[:limit]
return results
async def close(self):

View file

@ -1185,6 +1185,141 @@ async def test_filter_combinations(store: BaseVectorStore, _store_name: str):
logger.info("✓ Filter combinations test passed")
async def test_list_with_sorting(store: BaseVectorStore, _store_name: str):
"""Test list operation with sorting by timestamp to get most recent top 10 items."""
logger.info("=" * 20 + " LIST WITH SORTING TEST " + "=" * 20)
# Clean up any existing test data first
try:
existing_nodes = await store.list(filters={"test_type": "timestamp_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}")
# Create test nodes with timestamps
import time
test_nodes = []
base_timestamp = int(time.time())
for i in range(15):
node = VectorNode(
vector_id=f"timestamp_node_{i}",
content=f"Test content for node {i} with timestamp",
metadata={
"test_type": "timestamp_test",
"timestamp": base_timestamp - (14 - i) * 3600, # Each node is 1 hour newer
"index": i,
"created_at": base_timestamp - (14 - i) * 3600,
},
)
test_nodes.append(node)
# Insert test nodes
await store.insert(test_nodes)
logger.info(f"Inserted {len(test_nodes)} test nodes with timestamps")
# Test 1: Get all items sorted by timestamp (ascending)
results_asc = await store.list(
filters={"test_type": "timestamp_test"},
sort_key="timestamp",
reverse=False,
)
logger.info(f"Test 1 - Ascending order: {len(results_asc)} results")
# Verify ascending order
for i in range(len(results_asc) - 1):
ts1 = results_asc[i].metadata.get("timestamp", 0)
ts2 = results_asc[i + 1].metadata.get("timestamp", 0)
assert ts1 <= ts2, f"Results not in ascending order: {ts1} > {ts2}"
logger.info(
f"First item timestamp: {results_asc[0].metadata.get('timestamp')}, "
f"index: {results_asc[0].metadata.get('index')}",
)
logger.info(
f"Last item timestamp: {results_asc[-1].metadata.get('timestamp')}, "
f"index: {results_asc[-1].metadata.get('index')}",
)
# Test 2: Get top 10 most recent items (descending order)
results_desc = await store.list(
filters={"test_type": "timestamp_test"},
limit=10,
sort_key="timestamp",
reverse=True,
)
logger.info(f"Test 2 - Top 10 most recent (descending order): {len(results_desc)} results")
# Verify we got exactly 10 results
assert len(results_desc) == 10, f"Expected 10 results, got {len(results_desc)}"
# Log the actual results for debugging
logger.info("Top 10 results (should be index 14 to 5):")
for i, node in enumerate(results_desc):
logger.info(f" {i}: index={node.metadata.get('index')}, timestamp={node.metadata.get('timestamp')}")
# Verify descending order
for i in range(len(results_desc) - 1):
ts1 = results_desc[i].metadata.get("timestamp", 0)
ts2 = results_desc[i + 1].metadata.get("timestamp", 0)
assert ts1 >= ts2, f"Results not in descending order: {ts1} < {ts2}"
# Verify we got the most recent items (index 5-14)
for node in results_desc:
index = node.metadata.get("index", -1)
assert index >= 5, f"Top 10 should have index >= 5, got index {index}"
logger.info(
f"Most recent item - timestamp: {results_desc[0].metadata.get('timestamp')}, "
f"index: {results_desc[0].metadata.get('index')}",
)
logger.info(
f"10th most recent item - timestamp: {results_desc[-1].metadata.get('timestamp')}, "
f"index: {results_desc[-1].metadata.get('index')}",
)
# Test 3: Get top 5 most recent with additional filter
results_limited = await store.list(
filters={"test_type": "timestamp_test"},
limit=5,
sort_key="timestamp",
reverse=True,
)
logger.info(f"Test 3 - Top 5 most recent: {len(results_limited)} results")
assert len(results_limited) == 5, f"Expected 5 results, got {len(results_limited)}"
# Verify these are the 5 most recent
for i, node in enumerate(results_limited):
expected_index = 14 - i # Should be 14, 13, 12, 11, 10
actual_index = node.metadata.get("index", -1)
assert actual_index == expected_index, f"Expected index {expected_index}, got {actual_index}"
# Test 4: Sort by different key (created_at)
results_created = await store.list(
filters={"test_type": "timestamp_test"},
limit=3,
sort_key="created_at",
reverse=True,
)
logger.info(f"Test 4 - Top 3 by created_at: {len(results_created)} results")
assert len(results_created) == 3, f"Expected 3 results, got {len(results_created)}"
# Verify sorting by created_at
for i in range(len(results_created) - 1):
ts1 = results_created[i].metadata.get("created_at", 0)
ts2 = results_created[i + 1].metadata.get("created_at", 0)
assert ts1 >= ts2, f"Results not sorted by created_at: {ts1} < {ts2}"
# Clean up test data
await store.delete([node.vector_id for node in test_nodes])
logger.info("Cleaned up test nodes")
logger.info("✓ List with sorting test passed")
# ==================== Test Runner ====================
@ -1239,6 +1374,7 @@ async def run_all_tests_for_store(store_type: str, store_name: str):
await test_metadata_statistics(store, store_name)
await test_update_metadata_only(store, store_name)
await test_filter_combinations(store, store_name)
await test_list_with_sorting(store, store_name)
# ========== Collection Management Tests ==========
logger.info(f"\n{'#' * 60}")