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
18 changes: 8 additions & 10 deletions sc2/units.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,41 +50,39 @@ 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,
)

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,
)

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,
)

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,
)

Expand Down
8 changes: 8 additions & 0 deletions test/test_pickled_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading