Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 89 additions & 1 deletion test/collection/test_filter.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
import datetime
import uuid
from typing import Any, List

import pytest

Expand All @@ -12,7 +14,7 @@
_FilterValue,
_Operator,
)
from weaviate.collections.filters import _FilterToGRPC
from weaviate.collections.filters import _FilterToGRPC, _FilterToREST
from weaviate.proto.v1 import base_pb2


Expand Down Expand Up @@ -208,3 +210,89 @@ def test_reuse_by_ref_builder_for_independent_filters() -> None:
)
def test_operator_to_grpc(operator: _Operator, want: base_pb2.Filters.Operator) -> None:
assert operator._to_grpc() == want, "wrong pb operator"


# (values, gRPC array field, REST array key) for each supported filter value type
SEQUENCE_FILTER_VALUES = [
(["a", "b"], "value_text_array", "valueTextArray"),
([1, 2], "value_int_array", "valueIntArray"),
([1.5, 2.5], "value_number_array", "valueNumberArray"),
([True, False], "value_boolean_array", "valueBooleanArray"),
([uuid.UUID(int=1), uuid.UUID(int=2)], "value_text_array", "valueTextArray"),
(
[
datetime.datetime(2023, 1, 1, tzinfo=datetime.timezone.utc),
datetime.datetime(2024, 1, 1, tzinfo=datetime.timezone.utc),
],
"value_text_array",
"valueDateArray",
),
]
SEQUENCE_FILTER_IDS = ["text", "int", "float", "bool", "uuid", "date"]


@pytest.mark.parametrize("method", ["contains_any", "contains_all", "contains_none"])
@pytest.mark.parametrize(
"values,grpc_field,_rest_key", SEQUENCE_FILTER_VALUES, ids=SEQUENCE_FILTER_IDS
)
def test_sequence_filter_values_to_grpc(
method: str, values: List[Any], grpc_field: str, _rest_key: str
) -> None:
from_list = getattr(wvc.query.Filter.by_property("test"), method)(values)
from_tuple = getattr(wvc.query.Filter.by_property("test"), method)(tuple(values))

as_list = _FilterToGRPC.convert(from_list)
as_tuple = _FilterToGRPC.convert(from_tuple)

assert as_list.HasField(grpc_field), "the list form must send the values"
assert as_tuple == as_list, "a tuple must serialise like the equivalent list"


@pytest.mark.parametrize("method", ["contains_any", "contains_all", "contains_none"])
@pytest.mark.parametrize(
"values,_grpc_field,rest_key", SEQUENCE_FILTER_VALUES, ids=SEQUENCE_FILTER_IDS
)
def test_sequence_filter_values_to_rest(
method: str, values: List[Any], _grpc_field: str, rest_key: str
) -> None:
from_list = getattr(wvc.query.Filter.by_property("test"), method)(values)
from_tuple = getattr(wvc.query.Filter.by_property("test"), method)(tuple(values))

as_list = _FilterToREST.convert(from_list)

assert rest_key in as_list, "the list form must send the values"
assert _FilterToREST.convert(from_tuple) == as_list, (
"a tuple must serialise like the equivalent list"
)


@pytest.mark.parametrize("value", ["test", ""])
def test_string_filter_values_are_not_sequences(value: str) -> None:
filter_ = wvc.query.Filter.by_property("test").equal(value)

assert _FilterToGRPC.convert(filter_).value_text == value
assert _FilterToREST.convert(filter_) == {
"operator": "Equal",
"path": ["test"],
"valueText": value,
}


def test_empty_tuple_grpc_conversion() -> None:
"""Ensure the gRPC converter treats an empty tuple like an empty list."""
fv = _FilterValue(target="test", value=(), operator=_Operator.EQUAL)
with pytest.raises(weaviate.exceptions.WeaviateInvalidInputError):
_FilterToGRPC.convert(fv)


def test_empty_tuple_rest_conversion() -> None:
"""Ensure the REST converter treats an empty tuple like an empty list."""
fv = _FilterValue(target="test", value=(), operator=_Operator.EQUAL)
with pytest.raises(weaviate.exceptions.WeaviateInvalidInputError):
_FilterToREST.convert(fv)


@pytest.mark.parametrize("method", ["contains_any", "contains_all", "contains_none"])
def test_empty_tuple_input(method: str) -> None:
with pytest.raises(weaviate.exceptions.WeaviateInvalidInputError):
getattr(wvc.query.Filter.by_property("test"), method)(())
103 changes: 61 additions & 42 deletions weaviate/collections/filters.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import uuid as uuid_lib
from typing import Any, Dict, List, Literal, Optional, cast, overload
from typing import Any, Dict, List, Literal, Optional, Sequence, cast, overload

from weaviate.collections.classes.filters import (
FilterReturn,
Expand All @@ -20,6 +20,23 @@
from weaviate.util import _datetime_to_string


def _to_value_list(value: FilterValues) -> Optional[List[Any]]:
"""Return the filter value as a list if it holds multiple values, else `None`.

`FilterValuesList` is typed as a `Sequence`, so any sequence is valid user input.
`str` and `bytes` are sequences too, but are single filter values, not lists of them.

Args:
value: The filter value to normalise.

Returns:
The values as a list, or `None` if the filter value is a single value.
"""
if isinstance(value, Sequence) and not isinstance(value, (str, bytes)):
return list(value)
return None


class _FilterToGRPC:
@overload
@staticmethod
Expand All @@ -40,7 +57,8 @@ def convert(weav_filter: Optional[FilterReturn]) -> Optional[base_pb2.Filters]:

@staticmethod
def __value_filter(weav_filter: _FilterValue) -> base_pb2.Filters:
if isinstance(weav_filter.value, list) and len(weav_filter.value) == 0:
values = _to_value_list(weav_filter.value)
if values is not None and len(values) == 0:
raise WeaviateInvalidInputError(
"Filtering on empty lists is not supported by Weaviate. "
"To filter by property length, use "
Expand All @@ -64,10 +82,10 @@ def __value_filter(weav_filter: _FilterValue) -> base_pb2.Filters:
if isinstance(weav_filter.value, int) and not isinstance(weav_filter.value, bool)
else None,
value_number=(weav_filter.value if isinstance(weav_filter.value, float) else None),
value_boolean_array=_FilterToGRPC.__filter_to_bool_list(weav_filter.value),
value_int_array=_FilterToGRPC.__filter_to_int_list(weav_filter.value),
value_number_array=_FilterToGRPC.__filter_to_float_list(weav_filter.value),
value_text_array=_FilterToGRPC.__filter_to_text_list(weav_filter.value),
value_boolean_array=_FilterToGRPC.__filter_to_bool_list(values),
value_int_array=_FilterToGRPC.__filter_to_int_list(values),
value_number_array=_FilterToGRPC.__filter_to_float_list(values),
value_text_array=_FilterToGRPC.__filter_to_text_list(values),
value_geo=_FilterToGRPC.__filter_to_geo(weav_filter.value),
target=target,
)
Expand Down Expand Up @@ -121,52 +139,52 @@ def __filter_to_text(value: FilterValues) -> Optional[str]:
return _datetime_to_string(value)

@staticmethod
def __filter_to_text_list(value: FilterValues) -> Optional[base_pb2.TextArray]:
if not isinstance(value, list) or len(value) == 0:
def __filter_to_text_list(values: Optional[List[Any]]) -> Optional[base_pb2.TextArray]:
if values is None or len(values) == 0:
return None
if not (
isinstance(value[0], TIME)
or isinstance(value[0], str)
or isinstance(value[0], uuid_lib.UUID)
isinstance(values[0], TIME)
or isinstance(values[0], str)
or isinstance(values[0], uuid_lib.UUID)
):
return None

if isinstance(value[0], str):
value_list = value
elif isinstance(value[0], uuid_lib.UUID):
value_list = [str(uid) for uid in value]
if isinstance(values[0], str):
value_list = values
elif isinstance(values[0], uuid_lib.UUID):
value_list = [str(uid) for uid in values]
else:
dates = cast(List[TIME], value)
dates = cast(List[TIME], values)
value_list = [_datetime_to_string(date) for date in dates]

return base_pb2.TextArray(values=cast(List[str], value_list))

@staticmethod
def __filter_to_bool_list(value: FilterValues) -> Optional[base_pb2.BooleanArray]:
if not isinstance(value, list) or len(value) == 0 or not isinstance(value[0], bool):
def __filter_to_bool_list(values: Optional[List[Any]]) -> Optional[base_pb2.BooleanArray]:
if values is None or len(values) == 0 or not isinstance(values[0], bool):
return None

return base_pb2.BooleanArray(values=cast(List[bool], value))
return base_pb2.BooleanArray(values=cast(List[bool], values))

@staticmethod
def __filter_to_float_list(value: FilterValues) -> Optional[base_pb2.NumberArray]:
if not isinstance(value, list) or len(value) == 0 or not isinstance(value[0], float):
def __filter_to_float_list(values: Optional[List[Any]]) -> Optional[base_pb2.NumberArray]:
if values is None or len(values) == 0 or not isinstance(values[0], float):
return None

return base_pb2.NumberArray(values=cast(List[float], value))
return base_pb2.NumberArray(values=cast(List[float], values))

@staticmethod
def __filter_to_int_list(value: FilterValues) -> Optional[base_pb2.IntArray]:
def __filter_to_int_list(values: Optional[List[Any]]) -> Optional[base_pb2.IntArray]:
# bool is a subclass of int in Python, so the check must ensure it's not a bool
if (
not isinstance(value, list)
or len(value) == 0
or not isinstance(value[0], int)
or isinstance(value[0], bool)
values is None
or len(values) == 0
or not isinstance(values[0], int)
or isinstance(values[0], bool)
):
return None

return base_pb2.IntArray(values=cast(List[int], value))
return base_pb2.IntArray(values=cast(List[int], values))

@staticmethod
def __and_or_not_filter(weav_filter: FilterReturn) -> Optional[base_pb2.Filters]:
Expand Down Expand Up @@ -232,25 +250,26 @@ def __parse_filter(value: FilterValues) -> Dict[str, Any]:
return {"valueInt": value}
if isinstance(value, float):
return {"valueNumber": value}
if isinstance(value, list):
if len(value) == 0:
values = _to_value_list(value)
if values is not None:
if len(values) == 0:
raise WeaviateInvalidInputError(
"Filtering on empty lists is not supported by Weaviate. "
"To filter by property length, use "
"Filter.by_property('prop', length=True).equal(0)"
)
if isinstance(value[0], str):
return {"valueTextArray": value}
if isinstance(value[0], uuid_lib.UUID):
return {"valueTextArray": [str(val) for val in value]}
if isinstance(value[0], TIME):
return {"valueDateArray": [_datetime_to_string(cast(TIME, val)) for val in value]}
if isinstance(value[0], bool):
return {"valueBooleanArray": value}
if isinstance(value[0], int):
return {"valueIntArray": value}
if isinstance(value[0], float):
return {"valueNumberArray": value}
if isinstance(values[0], str):
return {"valueTextArray": values}
if isinstance(values[0], uuid_lib.UUID):
return {"valueTextArray": [str(val) for val in values]}
if isinstance(values[0], TIME):
return {"valueDateArray": [_datetime_to_string(cast(TIME, val)) for val in values]}
if isinstance(values[0], bool):
return {"valueBooleanArray": values}
if isinstance(values[0], int):
return {"valueIntArray": values}
if isinstance(values[0], float):
return {"valueNumberArray": values}
raise ValueError(f"Unknown filter value type: {type(value)}")

@staticmethod
Expand Down
Loading