diff --git a/docs/cookbook/appworld/quickstart.md b/docs/cookbook/appworld/quickstart.md index 0c18201a..78f6a1a8 100644 --- a/docs/cookbook/appworld/quickstart.md +++ b/docs/cookbook/appworld/quickstart.md @@ -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** diff --git a/reme_ai/core/config/default.yaml b/reme_ai/core/config/default.yaml index 08dfe23d..0b7276f7 100644 --- a/reme_ai/core/config/default.yaml +++ b/reme_ai/core/config/default.yaml @@ -30,7 +30,7 @@ embedding_model: vector_store: default: - backend: local + backend: chroma embedding_model: default collection_name: reme diff --git a/reme_ai/core/vector_store/base_vector_store.py b/reme_ai/core/vector_store/base_vector_store.py index 5815dca5..614666a4 100644 --- a/reme_ai/core/vector_store/base_vector_store.py +++ b/reme_ai/core/vector_store/base_vector_store.py @@ -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.""" diff --git a/reme_ai/core/vector_store/chroma_vector_store.py b/reme_ai/core/vector_store/chroma_vector_store.py index 231ca248..b639405f 100644 --- a/reme_ai/core/vector_store/chroma_vector_store.py +++ b/reme_ai/core/vector_store/chroma_vector_store.py @@ -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.""" diff --git a/reme_ai/core/vector_store/es_vector_store.py b/reme_ai/core/vector_store/es_vector_store.py index 0f1fa19f..6f52b3cc 100644 --- a/reme_ai/core/vector_store/es_vector_store.py +++ b/reme_ai/core/vector_store/es_vector_store.py @@ -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: diff --git a/reme_ai/core/vector_store/local_vector_store.py b/reme_ai/core/vector_store/local_vector_store.py index 04649841..a4f6d41b 100644 --- a/reme_ai/core/vector_store/local_vector_store.py +++ b/reme_ai/core/vector_store/local_vector_store.py @@ -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] diff --git a/reme_ai/core/vector_store/pgvector_store.py b/reme_ai/core/vector_store/pgvector_store.py index 695cee70..c02a0abd 100644 --- a/reme_ai/core/vector_store/pgvector_store.py +++ b/reme_ai/core/vector_store/pgvector_store.py @@ -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) diff --git a/reme_ai/core/vector_store/qdrant_vector_store.py b/reme_ai/core/vector_store/qdrant_vector_store.py index 8227c69b..fa48950a 100644 --- a/reme_ai/core/vector_store/qdrant_vector_store.py +++ b/reme_ai/core/vector_store/qdrant_vector_store.py @@ -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): diff --git a/tests/test_vector_store.py b/tests/test_vector_store.py index 76de21e1..2f4c1f4d 100644 --- a/tests/test_vector_store.py +++ b/tests/test_vector_store.py @@ -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}")