diff --git a/semantica/vector_store/milvus_store.py b/semantica/vector_store/milvus_store.py index 5271106a..9ac98faf 100644 --- a/semantica/vector_store/milvus_store.py +++ b/semantica/vector_store/milvus_store.py @@ -35,6 +35,7 @@ Author: Semantica Contributors License: MIT """ +import re from typing import Any, Dict, List, Optional, Union import numpy as np @@ -43,6 +44,40 @@ from ..utils.exceptions import ProcessingError, ValidationError from ..utils.logging import get_logger from ..utils.progress_tracker import get_progress_tracker + +def _validate_milvus_key(key: str) -> str: + """Validate and escape a metadata filter key for Milvus queries.""" + if not key or not isinstance(key, str) or not re.match(r"^[a-zA-Z0-9_.-]+$", key): + raise ValidationError(f"Invalid metadata filter key: '{key}'") + return key.replace("\\", "\\\\").replace('"', '\\"') + + +def _format_milvus_value(val: Any) -> str: + """Format and escape a filter value for Milvus expression syntax.""" + if isinstance(val, bool): + return "true" if val else "false" + elif isinstance(val, (int, float)): + return str(val) + elif isinstance(val, str): + escaped = ( + val.replace("\\", "\\\\") + .replace('"', '\\"') + .replace("\n", "\\n") + .replace("\r", "\\r") + ) + return f'"{escaped}"' + elif val is None: + return "null" + else: + escaped = ( + str(val) + .replace("\\", "\\\\") + .replace('"', '\\"') + .replace("\n", "\\n") + .replace("\r", "\\r") + ) + return f'"{escaped}"' + # Optional Milvus import try: from pymilvus import ( @@ -546,11 +581,11 @@ class MilvusStore: """Get vector by ID.""" if not MILVUS_AVAILABLE or not self.collection: return None - + try: - safe_id = vector_id.replace('"', '\\"') + safe_id = vector_id.replace("\\", "\\\\").replace('"', '\\"') res = self.collection.collection.query( - expr=f'id == "{safe_id}"', + expr=f'id == "{safe_id}"', output_fields=["vector"] ) if res and len(res) > 0: @@ -563,11 +598,11 @@ class MilvusStore: """Get metadata by ID.""" if not MILVUS_AVAILABLE or not self.collection: return None - + try: - safe_id = vector_id.replace('"', '\\"') + safe_id = vector_id.replace("\\", "\\\\").replace('"', '\\"') res = self.collection.collection.query( - expr=f'id == "{safe_id}"', + expr=f'id == "{safe_id}"', output_fields=["metadata"] ) if res and len(res) > 0: @@ -595,18 +630,22 @@ class MilvusStore: expr_parts = [] if filters: for key, value in filters.items(): + safe_key = _validate_milvus_key(key) if isinstance(value, dict): if "min" in value and value["min"] is not None: - expr_parts.append(f'metadata["{key}"] >= {value["min"]}') + min_val = _format_milvus_value(value["min"]) + expr_parts.append(f'metadata["{safe_key}"] >= {min_val}') if "max" in value and value["max"] is not None: - expr_parts.append(f'metadata["{key}"] <= {value["max"]}') + max_val = _format_milvus_value(value["max"]) + expr_parts.append(f'metadata["{safe_key}"] <= {max_val}') elif isinstance(value, list): - formatted_vals = [f'"{v}"' if isinstance(v, str) else str(v) for v in value] - expr_parts.append(f'metadata["{key}"] in [{", ".join(formatted_vals)}]') - elif isinstance(value, str): - expr_parts.append(f'metadata["{key}"] == "{value}"') + formatted_vals = [_format_milvus_value(v) for v in value] + expr_parts.append( + f'metadata["{safe_key}"] in [{", ".join(formatted_vals)}]' + ) else: - expr_parts.append(f'metadata["{key}"] == {value}') + formatted_val = _format_milvus_value(value) + expr_parts.append(f'metadata["{safe_key}"] == {formatted_val}') expr = " and ".join(expr_parts) if expr_parts else "id != ''" diff --git a/semantica/vector_store/pinecone_store.py b/semantica/vector_store/pinecone_store.py index 18ee7894..13c04d53 100644 --- a/semantica/vector_store/pinecone_store.py +++ b/semantica/vector_store/pinecone_store.py @@ -75,7 +75,7 @@ class PineconeClient: try: # Default to serverless spec if not provided - if spec is None: + if spec is None and ServerlessSpec is not None: spec = ServerlessSpec(cloud="aws", region="us-east-1") # Map metric names @@ -316,6 +316,7 @@ class PineconeStore: self.api_key = api_key or config.get("api_key") self.environment = environment or config.get("environment") + self.dimension: Optional[int] = config.get("dimension") self.client: Optional[PineconeClient] = None self.index: Optional[PineconeIndex] = None @@ -384,7 +385,7 @@ class PineconeStore: try: # Create index spec if not provided - if spec is None: + if spec is None and ServerlessSpec is not None: spec = ServerlessSpec(cloud="aws", region="us-east-1") self.client.create_index(index_name, dimension, metric, spec, **kwargs) @@ -393,6 +394,7 @@ class PineconeStore: pinecone_index = self.client.get_index(index_name) self.index = PineconeIndex(pinecone_index) self.search_engine = PineconeSearch(self.index) + self.dimension = dimension self.logger.info(f"Created Pinecone index: {index_name}") return self.index @@ -420,6 +422,13 @@ class PineconeStore: pinecone_index = self.client.get_index(index_name) self.index = PineconeIndex(pinecone_index) self.search_engine = PineconeSearch(self.index) + if self.dimension is None: + try: + stats = self.describe_index_stats() + if stats and isinstance(stats, dict) and stats.get("dimension"): + self.dimension = int(stats["dimension"]) + except Exception as e: + self.logger.warning(f"Could not determine index dimension for '{index_name}': {e}") return self.index except Exception as e: raise ProcessingError(f"Failed to get index: {str(e)}") @@ -505,6 +514,9 @@ class PineconeStore: else: vector_list.append(list(vector)) + if self.dimension is None and vector_list: + self.dimension = len(vector_list[0]) + self.progress_tracker.update_tracking( tracking_id, message="Upserting vectors to index..." ) @@ -571,6 +583,9 @@ class PineconeStore: else: query_vector = list(query_vector) + if self.dimension is None and query_vector: + self.dimension = len(query_vector) + results = self.search_engine.similarity_search( np.array(query_vector), k, filter, namespace, **options ) @@ -647,6 +662,22 @@ class PineconeStore: if self.index is None or not PINECONE_AVAILABLE: return [] + dimension = self.dimension + if dimension is None: + try: + stats = self.describe_index_stats() + if stats and isinstance(stats, dict) and stats.get("dimension"): + dimension = int(stats["dimension"]) + self.dimension = dimension + except Exception: + pass + + if not dimension: + raise ProcessingError( + "Index dimension is unknown. Please specify 'dimension' when initializing PineconeStore " + "or call create_index()/get_index() first." + ) + pinecone_filter = {} if filters: for key, value in filters.items(): @@ -663,7 +694,6 @@ class PineconeStore: else: pinecone_filter[key] = value - dimension = getattr(self, "dimension", 768) dummy_vector = [0.0] * dimension try: diff --git a/semantica/vector_store/weaviate_store.py b/semantica/vector_store/weaviate_store.py index 6ff480f3..b5551221 100644 --- a/semantica/vector_store/weaviate_store.py +++ b/semantica/vector_store/weaviate_store.py @@ -447,6 +447,49 @@ class WeaviateStore: self.logger.warning(f"Failed to get metadata for {vector_id}: {e}") return None + def _build_weaviate_filter(self, filters: Dict[str, Any]) -> Any: + """Build native Weaviate Filter object from metadata filter dictionary.""" + if not filters or not WEAVIATE_AVAILABLE: + return None + + Filter = None + try: + from weaviate.classes.query import Filter + except (ImportError, AttributeError): + try: + if weaviate and hasattr(weaviate, "classes") and hasattr(weaviate.classes, "query"): + Filter = getattr(weaviate.classes.query, "Filter", None) + except AttributeError: + Filter = None + + if Filter is None: + return None + + try: + conditions = [] + for key, value in filters.items(): + if isinstance(value, dict): + if "min" in value and value["min"] is not None: + conditions.append(Filter.by_property(key).greater_or_equal(value["min"])) + if "max" in value and value["max"] is not None: + conditions.append(Filter.by_property(key).less_or_equal(value["max"])) + elif isinstance(value, list): + conditions.append(Filter.by_property(key).contains_any(value)) + else: + conditions.append(Filter.by_property(key).equal(value)) + + if not conditions: + return None + + weaviate_filter = conditions[0] + for cond in conditions[1:]: + weaviate_filter = weaviate_filter & cond + + return weaviate_filter + except Exception as e: + self.logger.debug(f"Could not build native Weaviate filter: {e}") + return None + def filter_by_metadata( self, filters: Dict[str, Any], limit: int = 10 ) -> List[Dict[str, Any]]: @@ -465,28 +508,111 @@ class WeaviateStore: from .vector_store import _matches_filter + native_filter = self._build_weaviate_filter(filters) if filters else None + + results = [] + seen_ids = set() + after_cursor = None + scanned_count = 0 + page_size = max(limit, 100) + use_native_filter = native_filter is not None + try: - objs = self.collection.query.fetch_objects( - limit=limit, - include_vector=True - ) - results = [] - for obj in objs.objects: - properties = obj.properties or {} - if _matches_filter(properties, filters): - results.append( - { - "id": str(obj.uuid), - "metadata": properties, - "vector": np.array(obj.vector) if obj.vector else None, - } - ) - if len(results) >= limit: - break + while len(results) < limit: + kwargs = {"limit": page_size, "include_vector": True} + if use_native_filter and native_filter is not None: + kwargs["filters"] = native_filter + if after_cursor is not None: + kwargs["after"] = after_cursor + + try: + objs = self.collection.query.fetch_objects(**kwargs) + except TypeError as te: + # Handle kwargs incompatibility (e.g. mock or client version without filters/after) + if "filters" in kwargs: + use_native_filter = False + kwargs.pop("filters", None) + try: + objs = self.collection.query.fetch_objects(**kwargs) + except TypeError: + if "after" in kwargs: + kwargs.pop("after", None) + kwargs["offset"] = scanned_count + try: + objs = self.collection.query.fetch_objects(**kwargs) + except TypeError: + kwargs.pop("offset", None) + objs = self.collection.query.fetch_objects(**kwargs) + elif "after" in kwargs: + kwargs.pop("after", None) + kwargs["offset"] = scanned_count + try: + objs = self.collection.query.fetch_objects(**kwargs) + except TypeError: + kwargs.pop("offset", None) + objs = self.collection.query.fetch_objects(**kwargs) + else: + raise te + except Exception as fe: + if use_native_filter: + self.logger.warning( + f"Native Weaviate filter query failed, falling back to paginated fetch: {fe}" + ) + use_native_filter = False + kwargs.pop("filters", None) + objs = self.collection.query.fetch_objects(**kwargs) + else: + raise fe + + if not objs or not getattr(objs, "objects", None): + break + + batch_objects = objs.objects + if not batch_objects: + break + + new_objects_found = False + for obj in batch_objects: + obj_id = str(obj.uuid) if hasattr(obj, "uuid") and obj.uuid is not None else None + if obj_id: + if obj_id in seen_ids: + continue + seen_ids.add(obj_id) + new_objects_found = True + + properties = getattr(obj, "properties", None) or {} + if _matches_filter(properties, filters): + vector = None + if hasattr(obj, "vector") and obj.vector: + vector = np.array(obj.vector) + results.append( + { + "id": obj_id, + "metadata": properties, + "vector": vector, + } + ) + if len(results) >= limit: + break + + if not new_objects_found: + break + + scanned_count += len(batch_objects) + if len(batch_objects) < page_size: + break + + last_obj = batch_objects[-1] + if hasattr(last_obj, "uuid") and last_obj.uuid is not None: + after_cursor = str(last_obj.uuid) + else: + break + return results except Exception as e: self.logger.warning(f"Failed to fetch Weaviate objects by metadata filter: {e}") - return [] + return results if results else [] + def query_vectors( self, diff --git a/tests/vector_store/test_backend_metadata_filtering.py b/tests/vector_store/test_backend_metadata_filtering.py index fc5b1183..e86a5763 100644 --- a/tests/vector_store/test_backend_metadata_filtering.py +++ b/tests/vector_store/test_backend_metadata_filtering.py @@ -8,6 +8,7 @@ from semantica.vector_store.pinecone_store import PineconeStore from semantica.vector_store.milvus_store import MilvusStore from semantica.vector_store.pgvector_store import PgVectorStore from semantica.vector_store.weaviate_store import WeaviateStore +from semantica.utils.exceptions import ProcessingError, ValidationError class TestBackendMetadataFiltering(unittest.TestCase): @@ -52,7 +53,7 @@ class TestBackendMetadataFiltering(unittest.TestCase): @patch('semantica.vector_store.pinecone_store.PINECONE_AVAILABLE', True) def test_pinecone_store_filter_by_metadata(self): - store = PineconeStore() + store = PineconeStore(dimension=2) mock_index_wrapper = MagicMock() mock_inner_index = MagicMock() @@ -71,6 +72,19 @@ class TestBackendMetadataFiltering(unittest.TestCase): self.assertEqual(len(results), 1) self.assertEqual(results[0]["id"], "p1") self.assertEqual(results[0]["metadata"], {"status": "active"}) + # Assert query vector dimension matches store.dimension (2) + mock_inner_index.query.assert_called_once() + query_kw = mock_inner_index.query.call_args[1] + self.assertEqual(len(query_kw["vector"]), 2) + + @patch('semantica.vector_store.pinecone_store.PINECONE_AVAILABLE', True) + def test_pinecone_store_filter_by_metadata_unknown_dimension_raises(self): + store = PineconeStore() + mock_index_wrapper = MagicMock() + store.index = mock_index_wrapper + store.describe_index_stats = MagicMock(return_value={}) + with self.assertRaises(ProcessingError): + store.filter_by_metadata({"status": "active"}, limit=5) @patch('semantica.vector_store.milvus_store.MILVUS_AVAILABLE', True) def test_milvus_store_filter_by_metadata(self): @@ -88,6 +102,39 @@ class TestBackendMetadataFiltering(unittest.TestCase): self.assertEqual(results[0]["id"], "m1") self.assertEqual(results[0]["metadata"], {"lang": "py"}) + @patch('semantica.vector_store.milvus_store.MILVUS_AVAILABLE', True) + def test_milvus_store_filter_by_metadata_escaping(self): + store = MilvusStore() + mock_coll_wrapper = MagicMock() + mock_inner_coll = MagicMock() + mock_inner_coll.query.return_value = [] + mock_coll_wrapper.collection = mock_inner_coll + store.collection = mock_coll_wrapper + + store.filter_by_metadata( + { + "title": 'John "Jack" Doe', + "active": True, + "tags": ['python', 'c++ "v"'], + }, + limit=5, + ) + + mock_inner_coll.query.assert_called_once() + expr = mock_inner_coll.query.call_args[1]["expr"] + self.assertIn('metadata["title"] == "John \\"Jack\\" Doe"', expr) + self.assertIn('metadata["active"] == true', expr) + self.assertIn('metadata["tags"] in ["python", "c++ \\"v\\""]', expr) + + @patch('semantica.vector_store.milvus_store.MILVUS_AVAILABLE', True) + def test_milvus_store_filter_by_metadata_invalid_key_raises(self): + store = MilvusStore() + mock_coll_wrapper = MagicMock() + store.collection = mock_coll_wrapper + + with self.assertRaises(ValidationError): + store.filter_by_metadata({'dept" || 1==1 || "': "val"}, limit=5) + @patch('semantica.vector_store.pgvector_store.PSYCOPG3_AVAILABLE', True) @patch('semantica.vector_store.pgvector_store.psycopg_sql') def test_pgvector_store_filter_by_metadata(self, mock_sql): @@ -126,6 +173,85 @@ class TestBackendMetadataFiltering(unittest.TestCase): self.assertEqual(results[0]["id"], "w-uuid-1") self.assertEqual(results[0]["metadata"], {"dept": "eng"}) + def test_weaviate_store_filter_by_metadata_pagination(self): + """Test that WeaviateStore.filter_by_metadata paginates beyond page 1 to find matching items.""" + store = WeaviateStore() + mock_coll = MagicMock() + + # Batch 1: 100 non-matching objects + batch1_objs = [] + for i in range(100): + obj = MagicMock() + obj.uuid = f"batch1-uuid-{i}" + obj.properties = {"dept": "hr"} + obj.vector = [0.1, 0.1] + batch1_objs.append(obj) + + res1 = MagicMock() + res1.objects = batch1_objs + + # Batch 2: 2 matching objects + obj_match1 = MagicMock() + obj_match1.uuid = "match-uuid-1" + obj_match1.properties = {"dept": "eng"} + obj_match1.vector = [0.5, 0.5] + + obj_match2 = MagicMock() + obj_match2.uuid = "match-uuid-2" + obj_match2.properties = {"dept": "eng"} + obj_match2.vector = [0.6, 0.6] + + res2 = MagicMock() + res2.objects = [obj_match1, obj_match2] + + def side_effect(**kwargs): + if kwargs.get("after") == "batch1-uuid-99": + return res2 + return res1 + + mock_coll.query.fetch_objects.side_effect = side_effect + store.collection = mock_coll + + with patch('semantica.vector_store.weaviate_store.WEAVIATE_AVAILABLE', True): + results = store.filter_by_metadata({"dept": "eng"}, limit=5) + self.assertEqual(len(results), 2) + self.assertEqual(results[0]["id"], "match-uuid-1") + self.assertEqual(results[1]["id"], "match-uuid-2") + + def test_weaviate_store_filter_by_metadata_native_filter(self): + """Test building native Weaviate filters for exact, range, and list criteria.""" + store = WeaviateStore() + mock_filter_cls = MagicMock() + mock_filter_prop = MagicMock() + mock_filter_cls.by_property.return_value = mock_filter_prop + + mock_module = MagicMock() + mock_module.classes.query.Filter = mock_filter_cls + + with patch('semantica.vector_store.weaviate_store.WEAVIATE_AVAILABLE', True), \ + patch('semantica.vector_store.weaviate_store.weaviate', mock_module): + + # Test exact match + res = store._build_weaviate_filter({"dept": "eng"}) + mock_filter_cls.by_property.assert_called_with("dept") + mock_filter_prop.equal.assert_called_with("eng") + + # Test range filter + mock_filter_cls.reset_mock() + mock_filter_prop.reset_mock() + res = store._build_weaviate_filter({"age": {"min": 20, "max": 50}}) + mock_filter_cls.by_property.assert_called_with("age") + mock_filter_prop.greater_or_equal.assert_called_with(20) + mock_filter_prop.less_or_equal.assert_called_with(50) + + # Test list filter + mock_filter_cls.reset_mock() + mock_filter_prop.reset_mock() + res = store._build_weaviate_filter({"tags": ["a", "b"]}) + mock_filter_cls.by_property.assert_called_with("tags") + mock_filter_prop.contains_any.assert_called_with(["a", "b"]) + if __name__ == "__main__": unittest.main() + diff --git a/tests/vector_store/test_pinecone_store.py b/tests/vector_store/test_pinecone_store.py index f1a357b5..29fc0345 100644 --- a/tests/vector_store/test_pinecone_store.py +++ b/tests/vector_store/test_pinecone_store.py @@ -82,6 +82,7 @@ class TestPineconeStore(unittest.TestCase): self.assertIsInstance(store.search_engine, PineconeSearch) store.client.create_index.assert_called_once() + @patch('semantica.vector_store.pinecone_store.PINECONE_AVAILABLE', True) @patch('semantica.vector_store.pinecone_store.PineconeClientLib') def test_upsert_vectors(self, mock_pinecone_client): """Test upserting vectors to Pinecone index.""" @@ -105,6 +106,7 @@ class TestPineconeStore(unittest.TestCase): self.assertEqual(result["upserted_count"], 2) store.index.upsert_vectors.assert_called_once() + @patch('semantica.vector_store.pinecone_store.PINECONE_AVAILABLE', True) @patch('semantica.vector_store.pinecone_store.PineconeClientLib') def test_search_vectors(self, mock_pinecone_client): """Test searching vectors in Pinecone index.""" @@ -128,6 +130,7 @@ class TestPineconeStore(unittest.TestCase): self.assertEqual(results[0]["id"], "id1") store.search_engine.similarity_search.assert_called_once() + @patch('semantica.vector_store.pinecone_store.PINECONE_AVAILABLE', True) @patch('semantica.vector_store.pinecone_store.PineconeClientLib') def test_delete_vectors(self, mock_pinecone_client): """Test deleting vectors from Pinecone index.""" @@ -148,6 +151,7 @@ class TestPineconeStore(unittest.TestCase): # Fix: assert called without the empty dict store.index.delete_vectors.assert_called_once_with(["id1", "id2"], "") + @patch('semantica.vector_store.pinecone_store.PINECONE_AVAILABLE', True) @patch('semantica.vector_store.pinecone_store.PineconeClientLib') def test_fetch_vectors(self, mock_pinecone_client): """Test fetching vectors from Pinecone index."""