fixed/weavit_store,pinecone_store,milvus_store

This commit is contained in:
TaherTadpatri
2026-08-09 14:50:59 +05:30
parent b094268525
commit b6497ace41
5 changed files with 360 additions and 35 deletions
+52 -13
View File
@@ -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 != ''"
+33 -3
View File
@@ -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:
+144 -18
View File
@@ -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,
@@ -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()
@@ -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."""