From fdfda267190736fcb0e20eae0db387df2752fba7 Mon Sep 17 00:00:00 2001 From: Bas Schoenmaeckers Date: Mon, 1 Jun 2026 15:13:10 +0200 Subject: [PATCH] Add `set` object functions to c-api --- Cargo.lock | 1 + crates/capi/Cargo.toml | 1 + crates/capi/src/lib.rs | 1 + crates/capi/src/setobject.rs | 162 ++++++++++++++++++++++++++++++++++ crates/vm/src/builtins/set.rs | 10 +-- 5 files changed, 170 insertions(+), 5 deletions(-) create mode 100644 crates/capi/src/setobject.rs diff --git a/Cargo.lock b/Cargo.lock index d632ce7f741..e388302e489 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3479,6 +3479,7 @@ name = "rustpython-capi" version = "0.5.0" dependencies = [ "bitflags 2.11.1", + "itertools 0.14.0", "num-complex", "pyo3", "rustpython-stdlib", diff --git a/crates/capi/Cargo.toml b/crates/capi/Cargo.toml index 685b44d1ee0..70b7fc4b8c5 100644 --- a/crates/capi/Cargo.toml +++ b/crates/capi/Cargo.toml @@ -13,6 +13,7 @@ crate-type = ["cdylib", "rlib"] [dependencies] bitflags = { workspace = true } +itertools = { workspace = true } num-complex = { workspace = true } rustpython-vm = { workspace = true, features = ["threading", "compiler"] } rustpython-stdlib = {workspace = true, features = ["threading"] } diff --git a/crates/capi/src/lib.rs b/crates/capi/src/lib.rs index d7f9f8b7464..a2bc31b17bf 100644 --- a/crates/capi/src/lib.rs +++ b/crates/capi/src/lib.rs @@ -25,6 +25,7 @@ pub mod pyerrors; pub mod pylifecycle; pub mod pystate; pub mod refcount; +pub mod setobject; pub mod traceback; pub mod tupleobject; pub mod unicodeobject; diff --git a/crates/capi/src/setobject.rs b/crates/capi/src/setobject.rs new file mode 100644 index 00000000000..cc479371b27 --- /dev/null +++ b/crates/capi/src/setobject.rs @@ -0,0 +1,162 @@ +use crate::PyObject; +use crate::object::define_py_check; +use crate::pystate::with_vm; +use core::ffi::c_int; +use itertools::process_results; +use rustpython_vm::AsObject; +use rustpython_vm::PyPayload; +use rustpython_vm::TryFromObject; +use rustpython_vm::builtins::{PyFrozenSet, PySet}; +use rustpython_vm::function::ArgIterable; + +define_py_check!(fn PySet_Check, types.set_type); +define_py_check!(fn PyFrozenSet_Check, types.frozenset_type); + +#[unsafe(no_mangle)] +pub unsafe extern "C" fn PySet_New(iterable: *mut PyObject) -> *mut PyObject { + with_vm(|vm| { + if iterable.is_null() { + return Ok(PySet::default().into_ref(&vm.ctx)); + } + + let iterable = ArgIterable::try_from_object(vm, unsafe { &*iterable }.to_owned())?; + let set = PySet::default().into_ref(&vm.ctx); + for item in iterable.iter(vm)? { + set.add(item?, vm)?; + } + Ok(set) + }) +} + +#[unsafe(no_mangle)] +pub unsafe extern "C" fn PyFrozenSet_New(iterable: *mut PyObject) -> *mut PyObject { + with_vm(|vm| { + if iterable.is_null() { + return Ok(vm.ctx.empty_frozenset.to_owned()); + } + + let iterable = ArgIterable::try_from_object(vm, unsafe { &*iterable }.to_owned())?; + let set = process_results(iterable.iter(vm)?, |it| PyFrozenSet::from_iter(vm, it))??; + Ok(set.into_ref(&vm.ctx)) + }) +} + +#[unsafe(no_mangle)] +pub unsafe extern "C" fn PySet_Add(set: *mut PyObject, key: *mut PyObject) -> c_int { + with_vm(|vm| { + let set = unsafe { &*set }.try_downcast_ref::(vm)?; + let key = unsafe { &*key }.to_owned(); + set.add(key, vm) + }) +} + +#[unsafe(no_mangle)] +pub unsafe extern "C" fn PySet_Clear(set: *mut PyObject) -> c_int { + with_vm(|vm| { + let set = unsafe { &*set }.try_downcast_ref::(vm)?; + set.clear(); + Ok(()) + }) +} + +#[unsafe(no_mangle)] +pub unsafe extern "C" fn PySet_Contains(anyset: *mut PyObject, key: *mut PyObject) -> c_int { + with_vm(|vm| { + let anyset = unsafe { &*anyset }; + let key = unsafe { &*key }; + + if let Some(set) = anyset.downcast_ref::() { + set.__contains__(key, vm) + } else if let Some(frozenset) = anyset.downcast_ref::() { + frozenset.__contains__(key, vm) + } else { + Err(vm.new_type_error(format!( + "expected set or frozenset, got '{}'", + anyset.class().name() + ))) + } + }) +} + +#[unsafe(no_mangle)] +pub unsafe extern "C" fn PySet_Discard(set: *mut PyObject, key: *mut PyObject) -> c_int { + with_vm(|vm| { + let set = unsafe { &*set }.try_downcast_ref::(vm)?; + let key = unsafe { &*key }; + let had_item = set.__contains__(key, vm)?; + if had_item { + set.discard(key.to_owned(), vm)?; + } + Ok(had_item) + }) +} + +#[unsafe(no_mangle)] +pub unsafe extern "C" fn PySet_Pop(set: *mut PyObject) -> *mut PyObject { + with_vm(|vm| { + let set = unsafe { &*set }.try_downcast_ref::(vm)?; + set.pop(vm) + }) +} + +#[unsafe(no_mangle)] +pub unsafe extern "C" fn PySet_Size(anyset: *mut PyObject) -> isize { + with_vm(|vm| { + let anyset = unsafe { &*anyset }; + if let Some(set) = anyset.downcast_ref::() { + set.as_object().length(vm) + } else if let Some(frozenset) = anyset.downcast_ref::() { + frozenset.as_object().length(vm) + } else { + Err(vm.new_type_error(format!( + "expected set or frozenset, got '{}'", + anyset.class().name() + ))) + } + }) +} + +#[cfg(false)] +mod tests { + use pyo3::prelude::*; + use pyo3::types::{PyFrozenSet, PyInt, PySet}; + + #[test] + fn new_and_size() { + Python::attach(|py| { + let set = PySet::empty(py).unwrap(); + assert!(set.is_instance_of::()); + assert_eq!(set.len(), 0); + + let frozen = PyFrozenSet::empty(py).unwrap(); + assert!(frozen.is_instance_of::()); + assert_eq!(frozen.len(), 0); + }) + } + + #[test] + fn add_contains_discard() { + Python::attach(|py| { + let set = PySet::empty(py).unwrap(); + let item = PyInt::new(py, 42); + + set.add(&item).unwrap(); + assert!(set.contains(&item).unwrap()); + set.discard(&item).unwrap(); + assert!(!set.contains(&item).unwrap()); + }) + } + + #[test] + fn pop_reduces_size() { + Python::attach(|py| { + let set = PySet::empty(py).unwrap(); + set.add(7).unwrap(); + assert_eq!(set.len(), 1); + + let popped = set.pop().unwrap(); + assert_eq!(popped.extract::().unwrap(), 7); + assert_eq!(set.len(), 0); + }) + } +} diff --git a/crates/vm/src/builtins/set.rs b/crates/vm/src/builtins/set.rs index f730049ff01..136470c1813 100644 --- a/crates/vm/src/builtins/set.rs +++ b/crates/vm/src/builtins/set.rs @@ -537,7 +537,7 @@ impl PySet { self.inner.len() } - fn __contains__(&self, needle: &PyObject, vm: &VirtualMachine) -> PyResult { + pub fn __contains__(&self, needle: &PyObject, vm: &VirtualMachine) -> PyResult { self.inner.contains(needle, vm) } @@ -679,17 +679,17 @@ impl PySet { } #[pymethod] - fn discard(&self, item: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> { + pub fn discard(&self, item: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> { self.inner.discard(&item, vm).map(|_| ()) } #[pymethod] - fn clear(&self) { + pub fn clear(&self) { self.inner.clear() } #[pymethod] - fn pop(&self, vm: &VirtualMachine) -> PyResult { + pub fn pop(&self, vm: &VirtualMachine) -> PyResult { self.inner.pop(vm) } @@ -995,7 +995,7 @@ impl PyFrozenSet { self.inner.len() } - fn __contains__(&self, needle: &PyObject, vm: &VirtualMachine) -> PyResult { + pub fn __contains__(&self, needle: &PyObject, vm: &VirtualMachine) -> PyResult { self.inner.contains(needle, vm) }