Skip to main content

rustpython_vm/builtins/
coroutine.rs

1use super::{PyCode, PyGenericAlias, PyStrRef, PyTupleRef, PyType, PyTypeRef};
2use crate::{
3    AsObject, Context, Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult, VirtualMachine,
4    class::PyClassImpl,
5    coroutine::{Coro, warn_deprecated_throw_signature},
6    frame::FrameObjectRef,
7    function::OptionalArg,
8    object::{Traverse, TraverseFn},
9    protocol::PyIterReturn,
10    types::{Destructor, IterNext, Iterable, Representable, SelfIter},
11};
12use crossbeam_utils::atomic::AtomicCell;
13
14#[pyclass(module = false, name = "coroutine", traverse = "manual")]
15#[derive(Debug)]
16// PyCoro_Type in CPython
17pub struct PyCoroutine {
18    inner: Coro,
19    #[pymember(name = "cr_origin")]
20    origin: Option<PyTupleRef>,
21}
22
23unsafe impl Traverse for PyCoroutine {
24    fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
25        self.inner.traverse(tracer_fn);
26        if let Some(origin) = &self.origin {
27            origin.traverse(tracer_fn);
28        }
29    }
30}
31
32impl PyPayload for PyCoroutine {
33    // Tracked in `make_generator_or_coro`, together with the frame the object
34    // is born owning, so the pair costs one trip through the GC's gen0 list
35    // instead of two.
36    const NEW_REF_UNTRACKED: bool = true;
37
38    #[inline]
39    fn class(ctx: &Context) -> &'static Py<PyType> {
40        ctx.types.coroutine_type
41    }
42}
43
44impl PyCoroutine {
45    pub const fn as_coro(&self) -> &Coro {
46        &self.inner
47    }
48
49    #[must_use]
50    pub fn new(
51        frame: FrameObjectRef,
52        name: PyStrRef,
53        qualname: PyStrRef,
54        origin: Option<PyTupleRef>,
55    ) -> Self {
56        Self {
57            inner: Coro::new(frame, name, qualname),
58            origin,
59        }
60    }
61}
62
63#[pyclass(
64    itemsize = core::mem::size_of::<crate::PyObjectRef>(),
65    flags(DISALLOW_INSTANTIATION, HAS_WEAKREF),
66    with(Py, Representable, Destructor)
67)]
68impl PyCoroutine {}
69
70#[pyclass]
71impl Py<PyCoroutine> {
72    #[pymethod]
73    fn send(&self, value: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
74        self.inner.send(self.as_object(), value, vm)
75    }
76
77    #[pymethod]
78    fn throw(
79        &self,
80        exc_type: PyObjectRef,
81        exc_val: OptionalArg,
82        exc_tb: OptionalArg,
83        vm: &VirtualMachine,
84    ) -> PyResult<PyIterReturn> {
85        warn_deprecated_throw_signature(&exc_val, &exc_tb, vm)?;
86        self.inner.throw(
87            self.as_object(),
88            exc_type,
89            exc_val.unwrap_or_none(vm),
90            exc_tb.unwrap_or_none(vm),
91            vm,
92        )
93    }
94
95    #[pymethod]
96    fn close(&self, vm: &VirtualMachine) -> PyResult<PyObjectRef> {
97        self.inner.close(self.as_object(), vm)
98    }
99
100    #[pygetset]
101    // name of the coroutine
102    fn __name__(&self) -> PyStrRef {
103        self.inner.name()
104    }
105
106    #[pygetset(setter)]
107    fn set___name__(&self, name: PyStrRef) {
108        self.inner.set_name(name)
109    }
110
111    #[pygetset]
112    // qualified name of the coroutine
113    fn __qualname__(&self) -> PyStrRef {
114        self.inner.qualname()
115    }
116
117    #[pygetset(setter)]
118    fn set___qualname__(&self, qualname: PyStrRef) {
119        self.inner.set_qualname(qualname)
120    }
121
122    #[pymethod(name = "__await__")]
123    fn r#await(zelf: PyRef<PyCoroutine>) -> PyCoroutineWrapper {
124        PyCoroutineWrapper {
125            coro: zelf,
126            closed: AtomicCell::new(false),
127        }
128    }
129
130    #[pygetset]
131    fn cr_await(&self, _vm: &VirtualMachine) -> Option<PyObjectRef> {
132        self.inner.frame_opt().and_then(|f| f.yield_from_target())
133    }
134    #[pygetset]
135    fn cr_frame(&self, _vm: &VirtualMachine) -> Option<FrameObjectRef> {
136        if self.inner.closed() {
137            None
138        } else {
139            self.inner.frame_opt()
140        }
141    }
142    #[pygetset]
143    fn cr_running(&self, _vm: &VirtualMachine) -> bool {
144        self.inner.running()
145    }
146    #[pygetset]
147    fn cr_code(&self, _vm: &VirtualMachine) -> PyRef<PyCode> {
148        self.inner.code()
149    }
150    #[pygetset]
151    fn cr_suspended(&self, _vm: &VirtualMachine) -> bool {
152        self.inner.suspended()
153    }
154
155    #[pyclassmethod]
156    fn __class_getitem__(
157        cls: PyTypeRef,
158        args: PyObjectRef,
159        vm: &VirtualMachine,
160    ) -> PyResult<PyGenericAlias> {
161        PyGenericAlias::from_args(cls, args, vm)
162    }
163}
164
165impl Representable for PyCoroutine {
166    #[inline]
167    fn repr_str(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<String> {
168        Ok(zelf.inner.repr(zelf.as_object(), zelf.get_id(), vm))
169    }
170}
171
172impl Destructor for PyCoroutine {
173    fn del(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<()> {
174        if zelf.inner.closed() || zelf.inner.running() {
175            return Ok(());
176        }
177        if zelf.inner.frame_opt().is_none_or(|f| f.lasti() == 0) {
178            crate::warn::warn_unawaited_coroutine(zelf.as_object(), &zelf.inner.qualname(), vm);
179            zelf.inner.closed.store(true);
180            return Ok(());
181        }
182        if let Err(e) = zelf.inner.close(zelf.as_object(), vm) {
183            crate::coroutine::unraisable_while_closing(zelf.as_object(), &zelf.inner, e, vm);
184        }
185        Ok(())
186    }
187}
188
189#[pyclass(module = false, name = "coroutine_wrapper", traverse = "manual")]
190#[derive(Debug)]
191// PyCoroWrapper_Type in CPython
192pub(crate) struct PyCoroutineWrapper {
193    coro: PyRef<PyCoroutine>,
194    closed: AtomicCell<bool>,
195}
196
197unsafe impl Traverse for PyCoroutineWrapper {
198    fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
199        self.coro.traverse(tracer_fn);
200    }
201}
202
203impl PyPayload for PyCoroutineWrapper {
204    #[inline]
205    fn class(ctx: &Context) -> &'static Py<PyType> {
206        ctx.types.coroutine_wrapper_type
207    }
208}
209
210impl PyCoroutineWrapper {
211    fn check_closed(&self, vm: &VirtualMachine) -> PyResult<()> {
212        if self.closed.load() {
213            return Err(vm.new_runtime_error("cannot reuse already awaited coroutine"));
214        }
215        Ok(())
216    }
217}
218
219#[pyclass(with(IterNext, Iterable))]
220impl Py<PyCoroutineWrapper> {
221    #[pymethod]
222    fn send(&self, val: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
223        self.check_closed(vm)?;
224        let result = self.coro.send(val, vm);
225        // Mark as closed if exhausted
226        if let Ok(PyIterReturn::StopIteration(_)) = &result {
227            self.closed.store(true);
228        }
229        result
230    }
231
232    #[pymethod]
233    fn throw(
234        &self,
235        exc_type: PyObjectRef,
236        exc_val: OptionalArg,
237        exc_tb: OptionalArg,
238        vm: &VirtualMachine,
239    ) -> PyResult<PyIterReturn> {
240        self.check_closed(vm)?;
241        warn_deprecated_throw_signature(&exc_val, &exc_tb, vm)?;
242        let result = self.coro.throw(exc_type, exc_val, exc_tb, vm);
243        // Mark as closed if exhausted
244        if let Ok(PyIterReturn::StopIteration(_)) = &result {
245            self.closed.store(true);
246        }
247        result
248    }
249
250    #[pymethod]
251    fn close(&self, vm: &VirtualMachine) -> PyResult<PyObjectRef> {
252        self.closed.store(true);
253        self.coro.close(vm)
254    }
255}
256
257impl SelfIter for PyCoroutineWrapper {}
258impl IterNext for PyCoroutineWrapper {
259    fn next(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
260        zelf.send(vm.ctx.none(), vm)
261    }
262}
263
264impl Drop for PyCoroutine {
265    fn drop(&mut self) {
266        if let Some(frame) = self.inner.frame_opt() {
267            frame.clear_generator();
268        }
269    }
270}
271
272/// Fast, VM-free check mirroring the read-only-state branches of
273/// `<PyCoroutine as Destructor>::del`: an already-closed or currently
274/// running coroutine needs no `close()`-style cleanup, so `del` is a
275/// documented no-op. Skipping the call avoids attaching to a VM (`with_vm`)
276/// on every coroutine drop for the common case of a coroutine driven to
277/// completion (e.g. `await`ed to a `return`).
278fn coroutine_del_needed(zelf: &PyObject) -> bool {
279    let zelf: &Py<PyCoroutine> = zelf
280        .downcast_ref()
281        .expect("del_needed is only installed on the coroutine type");
282    !(zelf.inner.closed() || zelf.inner.running())
283}
284
285pub(crate) fn init(ctx: &'static Context) {
286    PyCoroutine::extend_class(ctx, ctx.types.coroutine_type);
287    PyCoroutineWrapper::extend_class(ctx, ctx.types.coroutine_wrapper_type);
288    ctx.types
289        .coroutine_type
290        .slots
291        .del_needed
292        .store(Some(coroutine_del_needed));
293}