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
25 changes: 14 additions & 11 deletions pyiceberg/expressions/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -698,19 +698,18 @@ class SetPredicate(UnboundPredicate, ABC):
def __init__(
self, term: str | UnboundTerm, literals: Iterable[Any] | Iterable[LiteralValue] | None = None, **kwargs: Any
) -> None:
if literals is None and "values" in kwargs:
literals = kwargs["values"]

if literals is None:
literal_set: set[LiteralValue] = set()
else:
literal_set = _to_literal_set(literals)
# __new__ already built the set, and may have used up a one-shot iterator doing it.
literal_set = self.__dict__.pop("_literals_from_new", None)
if literal_set is None:
if literals is None and "values" in kwargs:
literals = kwargs["values"]
literal_set = set() if literals is None else _to_literal_set(literals)
super().__init__(term=_to_unbound_term(term), values=literal_set)

def bind(self, schema: Schema, case_sensitive: bool = True) -> BoundSetPredicate:
bound_term = self.term.bind(schema, case_sensitive)
literal_set = self.literals
return self.as_bound(bound_term, {lit.to(bound_term.ref().field.field_type) for lit in literal_set}) # type: ignore
field_type = bound_term.ref().field.field_type
return self.as_bound(bound_term, {lit.to(field_type) for lit in self.literals}) # type: ignore

def __str__(self) -> str:
"""Return the string representation of the SetPredicate class."""
Expand Down Expand Up @@ -844,7 +843,9 @@ def __new__( # pylint: disable=W0221
elif count == 1:
return EqualTo(term, next(iter(literals_set)))
else:
return super().__new__(cls)
predicate = super().__new__(cls)
object.__setattr__(predicate, "_literals_from_new", literals_set)
return predicate

def __invert__(self) -> NotIn:
"""Transform the Expression into its negated version."""
Expand Down Expand Up @@ -879,7 +880,9 @@ def __new__( # pylint: disable=W0221
elif count == 1:
return NotEqualTo(term, next(iter(literals_set)))
else:
return super().__new__(cls)
predicate = super().__new__(cls)
object.__setattr__(predicate, "_literals_from_new", literals_set)
return predicate

def __invert__(self) -> In:
"""Transform the Expression into its negated version."""
Expand Down
8 changes: 8 additions & 0 deletions tests/expressions/test_expressions.py
Original file line number Diff line number Diff line change
Expand Up @@ -327,6 +327,14 @@ def test_in_list() -> None:
assert In(Reference("foo"), ["a", "bc", "def"]).literals == {literal("a"), literal("bc"), literal("def")}


def test_in_generator() -> None:
assert In(Reference("foo"), (v for v in ["a", "bc", "def"])).literals == {literal("a"), literal("bc"), literal("def")}


def test_not_in_generator() -> None:
assert NotIn(Reference("foo"), (v for v in ["a", "bc", "def"])).literals == {literal("a"), literal("bc"), literal("def")}


def test_not_in_empty() -> None:
assert NotIn(Reference("foo"), ()) == AlwaysTrue()

Expand Down
Loading