Skip to content

Commit f1ae6b4

Browse files
committed
Implemented set.intersection and set.intersection_update
1 parent 032129f commit f1ae6b4

2 files changed

Lines changed: 51 additions & 0 deletions

File tree

py/objset.c

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -175,6 +175,43 @@ static mp_obj_t set_diff_update(int n_args, const mp_obj_t *args) {
175175
}
176176
static MP_DEFINE_CONST_FUN_OBJ_VAR(set_diff_update_obj, 1, set_diff_update);
177177

178+
static mp_obj_t set_intersect_int(mp_obj_t self_in, mp_obj_t other, bool update) {
179+
assert(MP_OBJ_IS_TYPE(self_in, &set_type));
180+
if (self_in == other) {
181+
return update ? mp_const_none : set_copy(self_in);
182+
}
183+
184+
mp_obj_set_t *self = self_in;
185+
mp_obj_set_t *out = mp_obj_new_set(0, NULL);
186+
187+
mp_obj_t iter = rt_getiter(other);
188+
mp_obj_t next;
189+
while ((next = rt_iternext(iter)) != mp_const_stop_iteration) {
190+
if (mp_set_lookup(&self->set, next, MP_MAP_LOOKUP)) {
191+
set_add(out, next);
192+
}
193+
}
194+
195+
if (update) {
196+
m_del(mp_obj_t, self->set.table, self->set.alloc);
197+
self->set.alloc = out->set.alloc;
198+
self->set.used = out->set.used;
199+
self->set.table = out->set.table;
200+
}
201+
202+
return update ? mp_const_none : out;
203+
}
204+
205+
static mp_obj_t set_intersect(mp_obj_t self_in, mp_obj_t other) {
206+
return set_intersect_int(self_in, other, false);
207+
}
208+
static MP_DEFINE_CONST_FUN_OBJ_2(set_intersect_obj, set_intersect);
209+
210+
static mp_obj_t set_intersect_update(mp_obj_t self_in, mp_obj_t other) {
211+
return set_intersect_int(self_in, other, true);
212+
}
213+
static MP_DEFINE_CONST_FUN_OBJ_2(set_intersect_update_obj, set_intersect_update);
214+
178215

179216
/******************************************************************************/
180217
/* set constructors & public C API */
@@ -187,6 +224,8 @@ static const mp_method_t set_type_methods[] = {
187224
{ "discard", &set_discard_obj },
188225
{ "difference", &set_diff_obj },
189226
{ "difference_update", &set_diff_update_obj },
227+
{ "intersection", &set_intersect_obj },
228+
{ "intersection_update", &set_intersect_update_obj },
190229
{ NULL, NULL }, // end-of-list sentinel
191230
};
192231

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,12 @@
1+
def report(s):
2+
l = list(s)
3+
l.sort()
4+
print(l)
5+
6+
s = {1, 2, 3, 4}
7+
report(s)
8+
report(s.intersection({1, 3}))
9+
report(s.intersection([3, 4]))
10+
11+
print(s.intersection_update([1]))
12+
report(s)

0 commit comments

Comments
 (0)