diff --git a/sc2/units.py b/sc2/units.py index 3e58b1d2..7431db13 100644 --- a/sc2/units.py +++ b/sc2/units.py @@ -50,11 +50,9 @@ def __or__(self, other: Units) -> Units: """ :param other: """ + self_tags = {self_unit.tag for self_unit in self} return Units( - chain( - iter(self), - (other_unit for other_unit in other if other_unit.tag not in (self_unit.tag for self_unit in self)), - ), + chain(iter(self), (other_unit for other_unit in other if other_unit.tag not in self_tags)), self._bot_object, ) @@ -62,11 +60,9 @@ def __add__(self, other: Units) -> Units: """ :param other: """ + self_tags = {self_unit.tag for self_unit in self} return Units( - chain( - iter(self), - (other_unit for other_unit in other if other_unit.tag not in (self_unit.tag for self_unit in self)), - ), + chain(iter(self), (other_unit for other_unit in other if other_unit.tag not in self_tags)), self._bot_object, ) @@ -74,8 +70,9 @@ def __and__(self, other: Units) -> Units: """ :param other: """ + self_tags = {self_unit.tag for self_unit in self} return Units( - (other_unit for other_unit in other if other_unit.tag in (self_unit.tag for self_unit in self)), + (other_unit for other_unit in other if other_unit.tag in self_tags), self._bot_object, ) @@ -83,8 +80,9 @@ def __sub__(self, other: Units) -> Units: """ :param other: """ + other_tags = {other_unit.tag for other_unit in other} return Units( - (self_unit for self_unit in self if self_unit.tag not in (other_unit.tag for other_unit in other)), + (self_unit for self_unit in self if self_unit.tag not in other_tags), self._bot_object, ) diff --git a/test/test_pickled_data.py b/test/test_pickled_data.py index fac77b71..a867f09f 100644 --- a/test/test_pickled_data.py +++ b/test/test_pickled_data.py @@ -804,6 +804,14 @@ def test_units(): assert scvs.random_or(1) assert townhalls.random_or(0) assert scvs.random_group_of(11) + # Set operations dedupe by tag and keep the order: self first, then other + head = Units(scvs[:6], bot) + tail = Units(scvs[4:], bot) + assert [u.tag for u in head | tail] == [u.tag for u in scvs] + assert [u.tag for u in head + tail] == [u.tag for u in scvs] + assert [u.tag for u in head & tail] == [u.tag for u in scvs[4:6]] + assert [u.tag for u in head - tail] == [u.tag for u in scvs[:4]] + assert [u.tag for u in head - head] == [] assert not scvs.random_group_of(0) assert not townhalls.random_group_of(0) # assert not scvs.in_attack_range_of(townhalls.first)