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)]
16pub 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 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 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 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)]
191pub(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 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 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
272fn 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}