mirror of
https://github.com/semantica-agi/semantica.git
synced 2026-08-29 04:26:20 +00:00
fixed/weavit_store,pinecone_store,milvus_store
This commit is contained in:
@@ -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 != ''"
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user