mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-11 03:40:03 +00:00
feat(vector-store): add sorting capability to list operations across all backends
This commit is contained in:
parent
2d98169d80
commit
a172a6854b
9 changed files with 269 additions and 12 deletions
|
|
@ -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**
|
||||
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ embedding_model:
|
|||
|
||||
vector_store:
|
||||
default:
|
||||
backend: local
|
||||
backend: chroma
|
||||
embedding_model: default
|
||||
collection_name: reme
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue