1use super::{PyCode, PyGenerator, PyGenericAlias, PyStrRef, PyType, PyTypeRef};
2use crate::{
3 AsObject, Context, Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult, VirtualMachine,
4 builtins::PyBaseExceptionRef,
5 class::PyClassImpl,
6 common::lock::PyMutex,
7 coroutine::{Coro, warn_deprecated_throw_signature},
8 frame::FrameObjectRef,
9 function::OptionalArg,
10 object::{Traverse, TraverseFn},
11 protocol::PyIterReturn,
12 types::{Destructor, IterNext, Iterable, Representable, SelfIter},
13};
14
15use core::sync::atomic::{AtomicBool, Ordering};
16use crossbeam_utils::atomic::AtomicCell;
17
18fn warn_unawaited_asyncgen_method(ag: &Py<PyAsyncGen>, method: &str, vm: &VirtualMachine) {
19 let name = ag.as_coro().qualname();
20 let msg = format!("coroutine method '{method}' of '{name}' was never awaited");
21 if let Err(e) = crate::stdlib::_warnings::warn(vm.ctx.exceptions.runtime_warning, msg, 1, vm) {
22 vm.run_unraisable(e, None, ag.as_object().to_owned());
23 }
24}
25
26#[pyclass(name = "async_generator", module = false, traverse = "manual")]
27#[derive(Debug)]
28pub struct PyAsyncGen {
29 inner: Coro,
30 #[pymember(name = "ag_running")]
31 running_async: AtomicBool,
32 ag_hooks_inited: AtomicCell<bool>,
34 ag_closed: AtomicCell<bool>,
38 ag_finalizer: PyMutex<Option<PyObjectRef>>,
40}
41
42unsafe impl Traverse for PyAsyncGen {
43 fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
44 self.inner.traverse(tracer_fn);
45 self.ag_finalizer.traverse(tracer_fn);
46 }
47}
48type PyAsyncGenRef = PyRef<PyAsyncGen>;
49
50impl PyPayload for PyAsyncGen {
51 const NEW_REF_UNTRACKED: bool = true;
55
56 #[inline]
57 fn class(ctx: &Context) -> &'static Py<PyType> {
58 ctx.types.async_generator
59 }
60}
61
62impl PyAsyncGen {
63 pub const fn as_coro(&self) -> &Coro {
64 &self.inner
65 }
66
67 #[must_use]
68 pub fn new(frame: FrameObjectRef, name: PyStrRef, qualname: PyStrRef) -> Self {
69 Self {
70 inner: Coro::new(frame, name, qualname),
71 running_async: AtomicBool::new(false),
72 ag_hooks_inited: AtomicCell::new(false),
73 ag_closed: AtomicCell::new(false),
74 ag_finalizer: PyMutex::new(None),
75 }
76 }
77
78 fn init_hooks(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<()> {
81 if zelf.ag_hooks_inited.load() {
83 return Ok(());
84 }
85
86 zelf.ag_hooks_inited.store(true);
87
88 let finalizer = vm.async_gen_finalizer.borrow().clone();
90 if let Some(finalizer) = finalizer {
91 *zelf.ag_finalizer.lock() = Some(finalizer);
92 }
93
94 let firstiter = vm.async_gen_firstiter.borrow().clone();
96 if let Some(firstiter) = firstiter {
97 let obj: PyObjectRef = zelf.to_owned().into();
98 firstiter.call((obj,), vm)?;
99 }
100
101 Ok(())
102 }
103
104 fn call_finalizer(zelf: &Py<Self>, vm: &VirtualMachine) {
106 let finalizer = zelf.ag_finalizer.lock().clone();
107 if let Some(finalizer) = finalizer
108 && !zelf.ag_closed.load()
109 {
110 let obj: PyObjectRef = zelf.to_owned().into();
113
114 if let Err(e) = finalizer.call((obj,), vm) {
116 vm.run_unraisable(e, Some("async generator finalizer".to_owned()), finalizer);
117 }
118 }
119 }
120}
121
122#[pyclass(
123 itemsize = core::mem::size_of::<crate::PyObjectRef>(),
124 flags(DISALLOW_INSTANTIATION, HAS_WEAKREF),
125 with(PyRef, Representable, Destructor)
126)]
127impl Py<PyAsyncGen> {
128 #[pygetset]
129 fn __name__(&self) -> PyStrRef {
130 self.inner.name()
131 }
132
133 #[pygetset(setter)]
134 fn set___name__(&self, name: PyStrRef) {
135 self.inner.set_name(name)
136 }
137
138 #[pygetset]
139 fn __qualname__(&self) -> PyStrRef {
140 self.inner.qualname()
141 }
142
143 #[pygetset(setter)]
144 fn set___qualname__(&self, qualname: PyStrRef) {
145 self.inner.set_qualname(qualname)
146 }
147
148 #[pygetset]
149 fn ag_await(&self, _vm: &VirtualMachine) -> Option<PyObjectRef> {
150 self.inner.frame_opt().and_then(|f| f.yield_from_target())
151 }
152 #[pygetset]
153 fn ag_frame(&self, _vm: &VirtualMachine) -> Option<FrameObjectRef> {
154 if self.inner.closed() {
155 None
156 } else {
157 self.inner.frame_opt()
158 }
159 }
160 #[pygetset]
161 fn ag_code(&self, _vm: &VirtualMachine) -> PyRef<PyCode> {
162 self.inner.code()
163 }
164 #[pygetset]
165 fn ag_suspended(&self, _vm: &VirtualMachine) -> bool {
166 self.inner.suspended()
167 }
168
169 #[pyclassmethod]
170 fn __class_getitem__(
171 cls: PyTypeRef,
172 args: PyObjectRef,
173 vm: &VirtualMachine,
174 ) -> PyResult<PyGenericAlias> {
175 PyGenericAlias::from_args(cls, args, vm)
176 }
177}
178
179#[pyclass]
180impl PyRef<PyAsyncGen> {
181 #[pymethod]
182 const fn __aiter__(self, _vm: &VirtualMachine) -> Self {
183 self
184 }
185
186 #[pymethod]
187 fn __anext__(self, vm: &VirtualMachine) -> PyResult<PyAsyncGenASend> {
188 PyAsyncGen::init_hooks(&self, vm)?;
189 Ok(PyAsyncGenASend {
190 ag: self,
191 state: AtomicCell::new(AwaitableState::Init),
192 value: vm.ctx.none(),
193 })
194 }
195
196 #[pymethod]
197 fn asend(self, value: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyAsyncGenASend> {
198 PyAsyncGen::init_hooks(&self, vm)?;
199 Ok(PyAsyncGenASend {
200 ag: self,
201 state: AtomicCell::new(AwaitableState::Init),
202 value,
203 })
204 }
205
206 #[pymethod]
207 fn athrow(
208 self,
209 exc_type: PyObjectRef,
210 exc_val: OptionalArg,
211 exc_tb: OptionalArg,
212 vm: &VirtualMachine,
213 ) -> PyResult<PyAsyncGenAThrow> {
214 warn_deprecated_throw_signature(&exc_val, &exc_tb, vm)?;
215 PyAsyncGen::init_hooks(&self, vm)?;
216 Ok(PyAsyncGenAThrow {
217 ag: self,
218 aclose: false,
219 state: AtomicCell::new(AwaitableState::Init),
220 value: (
221 exc_type,
222 exc_val.unwrap_or_none(vm),
223 exc_tb.unwrap_or_none(vm),
224 ),
225 })
226 }
227
228 #[pymethod]
229 fn aclose(self, vm: &VirtualMachine) -> PyResult<PyAsyncGenAThrow> {
230 PyAsyncGen::init_hooks(&self, vm)?;
231 Ok(PyAsyncGenAThrow {
232 ag: self,
233 aclose: true,
234 state: AtomicCell::new(AwaitableState::Init),
235 value: (
236 vm.ctx.exceptions.generator_exit.to_owned().into(),
237 vm.ctx.none(),
238 vm.ctx.none(),
239 ),
240 })
241 }
242}
243
244impl Representable for PyAsyncGen {
245 #[inline]
246 fn repr_str(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<String> {
247 Ok(zelf.inner.repr(zelf.as_object(), zelf.get_id(), vm))
248 }
249}
250
251#[pyclass(
252 module = false,
253 name = "async_generator_wrapped_value",
254 traverse = "manual"
255)]
256#[derive(Debug)]
257pub(crate) struct PyAsyncGenWrappedValue(pub PyObjectRef);
258
259unsafe impl Traverse for PyAsyncGenWrappedValue {
260 fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
261 self.0.traverse(tracer_fn);
262 }
263}
264
265impl PyPayload for PyAsyncGenWrappedValue {
266 #[inline]
267 fn class(ctx: &Context) -> &'static Py<PyType> {
268 ctx.types.async_generator_wrapped_value
269 }
270}
271
272#[pyclass]
273impl PyAsyncGenWrappedValue {}
274
275impl PyAsyncGenWrappedValue {
276 fn unbox(ag: &Py<PyAsyncGen>, val: PyResult<PyIterReturn>, vm: &VirtualMachine) -> PyResult {
277 let (frame_done, mark_ag_closed, async_done) = match &val {
278 Ok(PyIterReturn::StopIteration(_)) => (true, true, true),
279 Err(e) if e.fast_isinstance(vm.ctx.exceptions.generator_exit) => (true, true, true),
280 Err(e) if e.fast_isinstance(vm.ctx.exceptions.stop_async_iteration) => {
281 (false, true, true)
282 }
283 Err(_) => (false, false, true),
284 _ => (false, false, false),
285 };
286 if frame_done {
287 ag.inner.closed.store(true);
288 }
289 if mark_ag_closed {
290 ag.ag_closed.store(true);
291 }
292 if async_done {
293 ag.running_async.store(false, Ordering::Relaxed);
294 }
295 let val = val?.into_async_pyresult(vm)?;
296 match_class!(match val {
297 val @ Self => {
298 ag.running_async.store(false, Ordering::Relaxed);
299 Err(vm.new_stop_iteration(Some(val.0.clone())))
300 }
301 val => Ok(val),
302 })
303 }
304}
305
306#[derive(Debug, Clone, Copy)]
307enum AwaitableState {
308 Init,
309 Iter,
310 Closed,
311}
312
313#[pyclass(module = false, name = "async_generator_asend", traverse = "manual")]
314#[derive(Debug)]
315pub(crate) struct PyAsyncGenASend {
316 ag: PyAsyncGenRef,
317 state: AtomicCell<AwaitableState>,
318 value: PyObjectRef,
319}
320
321unsafe impl Traverse for PyAsyncGenASend {
322 fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
323 self.ag.traverse(tracer_fn);
324 self.value.traverse(tracer_fn);
325 }
326}
327
328impl PyPayload for PyAsyncGenASend {
329 #[inline]
330 fn class(ctx: &Context) -> &'static Py<PyType> {
331 ctx.types.async_generator_asend
332 }
333}
334
335impl PyAsyncGenASend {
336 fn set_closed(&self) {
337 self.state.store(AwaitableState::Closed);
338 }
339}
340
341#[pyclass(with(IterNext, Iterable, Destructor))]
342impl Py<PyAsyncGenASend> {
343 #[pymethod(name = "__await__")]
344 const fn r#await(zelf: PyRef<PyAsyncGenASend>, _vm: &VirtualMachine) -> PyRef<PyAsyncGenASend> {
345 zelf
346 }
347
348 #[pymethod]
349 fn send(&self, val: PyObjectRef, vm: &VirtualMachine) -> PyResult {
350 let val = match self.state.load() {
351 AwaitableState::Closed => {
352 return Err(
353 vm.new_runtime_error("cannot reuse already awaited __anext__()/asend()")
354 );
355 }
356 AwaitableState::Iter => val, AwaitableState::Init => {
358 if self.ag.running_async.load(Ordering::Relaxed) {
359 self.state.store(AwaitableState::Closed);
360 return Err(
361 vm.new_runtime_error("anext(): asynchronous generator is already running")
362 );
363 }
364 self.ag.running_async.store(true, Ordering::Relaxed);
365 self.state.store(AwaitableState::Iter);
366 if vm.is_none(&val) {
367 self.value.clone()
368 } else {
369 val
370 }
371 }
372 };
373 let res = self.ag.inner.send(self.ag.as_object(), val, vm);
374 let res = PyAsyncGenWrappedValue::unbox(&self.ag, res, vm);
375 if res.is_err() {
376 self.set_closed();
377 }
378 res
379 }
380
381 #[pymethod]
382 fn throw(
383 &self,
384 exc_type: PyObjectRef,
385 exc_val: OptionalArg,
386 exc_tb: OptionalArg,
387 vm: &VirtualMachine,
388 ) -> PyResult {
389 match self.state.load() {
390 AwaitableState::Closed => {
391 return Err(
392 vm.new_runtime_error("cannot reuse already awaited __anext__()/asend()")
393 );
394 }
395 AwaitableState::Init => {
396 if self.ag.running_async.load(Ordering::Relaxed) {
397 self.state.store(AwaitableState::Closed);
398 return Err(
399 vm.new_runtime_error("anext(): asynchronous generator is already running")
400 );
401 }
402 self.ag.running_async.store(true, Ordering::Relaxed);
403 self.state.store(AwaitableState::Iter);
404 }
405 AwaitableState::Iter => {}
406 }
407
408 warn_deprecated_throw_signature(&exc_val, &exc_tb, vm)?;
409 let res = self.ag.inner.throw(
410 self.ag.as_object(),
411 exc_type,
412 exc_val.unwrap_or_none(vm),
413 exc_tb.unwrap_or_none(vm),
414 vm,
415 );
416 let res = PyAsyncGenWrappedValue::unbox(&self.ag, res, vm);
417 if res.is_err() {
418 self.set_closed();
419 }
420 res
421 }
422
423 #[pymethod]
424 fn close(&self, vm: &VirtualMachine) -> PyResult<()> {
425 if matches!(self.state.load(), AwaitableState::Closed) {
426 return Ok(());
427 }
428 let result = self.throw(
429 vm.ctx.exceptions.generator_exit.to_owned().into(),
430 OptionalArg::Missing,
431 OptionalArg::Missing,
432 vm,
433 );
434 match result {
435 Ok(_) => Err(vm.new_runtime_error("coroutine ignored GeneratorExit")),
436 Err(e)
437 if e.fast_isinstance(vm.ctx.exceptions.stop_iteration)
438 || e.fast_isinstance(vm.ctx.exceptions.stop_async_iteration)
439 || e.fast_isinstance(vm.ctx.exceptions.generator_exit) =>
440 {
441 Ok(())
442 }
443 Err(e) => Err(e),
444 }
445 }
446}
447
448impl SelfIter for PyAsyncGenASend {}
449impl IterNext for PyAsyncGenASend {
450 fn next(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
451 PyIterReturn::from_pyresult(zelf.send(vm.ctx.none(), vm), vm)
452 }
453}
454
455impl Destructor for PyAsyncGenASend {
456 fn del(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<()> {
457 if matches!(zelf.state.load(), AwaitableState::Init) {
458 warn_unawaited_asyncgen_method(&zelf.ag, "asend", vm);
459 }
460 Ok(())
461 }
462}
463
464#[pyclass(module = false, name = "async_generator_athrow", traverse = "manual")]
465#[derive(Debug)]
466pub(crate) struct PyAsyncGenAThrow {
467 ag: PyAsyncGenRef,
468 aclose: bool,
469 state: AtomicCell<AwaitableState>,
470 value: (PyObjectRef, PyObjectRef, PyObjectRef),
471}
472
473unsafe impl Traverse for PyAsyncGenAThrow {
474 fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
475 self.ag.traverse(tracer_fn);
476 self.value.traverse(tracer_fn);
477 }
478}
479
480impl PyPayload for PyAsyncGenAThrow {
481 #[inline]
482 fn class(ctx: &Context) -> &'static Py<PyType> {
483 ctx.types.async_generator_athrow
484 }
485}
486
487impl PyAsyncGenAThrow {
488 fn ignored_close(&self, res: &PyResult<PyIterReturn>) -> bool {
489 res.as_ref().is_ok_and(|v| match v {
490 PyIterReturn::Return(obj) => obj.downcastable::<PyAsyncGenWrappedValue>(),
491 PyIterReturn::StopIteration(_) => false,
492 })
493 }
494 fn yield_close(&self, vm: &VirtualMachine) -> PyBaseExceptionRef {
495 self.ag.running_async.store(false, Ordering::Relaxed);
496 self.state.store(AwaitableState::Closed);
497 vm.new_runtime_error("async generator ignored GeneratorExit")
498 }
499 fn check_error(&self, exc: PyBaseExceptionRef, vm: &VirtualMachine) -> PyBaseExceptionRef {
500 self.ag.running_async.store(false, Ordering::Relaxed);
501 self.state.store(AwaitableState::Closed);
502 if self.aclose
503 && (exc.fast_isinstance(vm.ctx.exceptions.stop_async_iteration)
504 || exc.fast_isinstance(vm.ctx.exceptions.generator_exit))
505 {
506 vm.new_stop_iteration(None)
507 } else {
508 exc
509 }
510 }
511}
512
513#[pyclass(with(IterNext, Iterable, Destructor))]
514impl Py<PyAsyncGenAThrow> {
515 #[pymethod(name = "__await__")]
516 const fn r#await(
517 zelf: PyRef<PyAsyncGenAThrow>,
518 _vm: &VirtualMachine,
519 ) -> PyRef<PyAsyncGenAThrow> {
520 zelf
521 }
522
523 #[pymethod]
524 fn send(&self, val: PyObjectRef, vm: &VirtualMachine) -> PyResult {
525 if matches!(self.state.load(), AwaitableState::Closed) {
526 return Err(vm.new_runtime_error("cannot reuse already awaited aclose()/athrow()"));
527 }
528 if self.ag.inner.closed() {
529 self.state.store(AwaitableState::Closed);
530 return Err(vm.new_stop_iteration(None));
531 }
532 if matches!(self.state.load(), AwaitableState::Init) {
533 if self.ag.running_async.load(Ordering::Relaxed) {
534 self.state.store(AwaitableState::Closed);
535 let msg = if self.aclose {
536 "aclose(): asynchronous generator is already running"
537 } else {
538 "athrow(): asynchronous generator is already running"
539 };
540 return Err(vm.new_runtime_error(msg.to_owned()));
541 }
542 if self.ag.ag_closed.load() {
543 self.state.store(AwaitableState::Closed);
544 return Err(
545 vm.new_exception_empty(vm.ctx.exceptions.stop_async_iteration.to_owned())
546 );
547 }
548 if !vm.is_none(&val) {
549 return Err(
550 vm.new_runtime_error("can't send non-None value to a just-started coroutine")
551 );
552 }
553 self.state.store(AwaitableState::Iter);
554 self.ag.running_async.store(true, Ordering::Relaxed);
555 if self.aclose {
556 self.ag.ag_closed.store(true);
557 }
558
559 let (ty, val, tb) = self.value.clone();
560 let ret = self.ag.inner.throw(self.ag.as_object(), ty, val, tb, vm);
561 if self.aclose && self.ignored_close(&ret) {
562 return Err(self.yield_close(vm));
563 }
564 let ret = if self.aclose {
565 ret.and_then(|o| o.into_async_pyresult(vm))
566 } else {
567 PyAsyncGenWrappedValue::unbox(&self.ag, ret, vm)
568 };
569 return ret.map_err(|e| self.check_error(e, vm));
570 }
571
572 let ret = self.ag.inner.send(self.ag.as_object(), val, vm);
573 if self.aclose {
574 match ret {
575 Ok(PyIterReturn::Return(v)) if v.downcastable::<PyAsyncGenWrappedValue>() => {
576 Err(self.yield_close(vm))
577 }
578 other => other
579 .and_then(|o| o.into_async_pyresult(vm))
580 .map_err(|e| self.check_error(e, vm)),
581 }
582 } else {
583 PyAsyncGenWrappedValue::unbox(&self.ag, ret, vm)
584 }
585 }
586
587 #[pymethod]
588 fn throw(
589 &self,
590 exc_type: PyObjectRef,
591 exc_val: OptionalArg,
592 exc_tb: OptionalArg,
593 vm: &VirtualMachine,
594 ) -> PyResult {
595 match self.state.load() {
596 AwaitableState::Closed => {
597 return Err(vm.new_runtime_error("cannot reuse already awaited aclose()/athrow()"));
598 }
599 AwaitableState::Init => {
600 if self.ag.running_async.load(Ordering::Relaxed) {
601 self.state.store(AwaitableState::Closed);
602 let msg = if self.aclose {
603 "aclose(): asynchronous generator is already running"
604 } else {
605 "athrow(): asynchronous generator is already running"
606 };
607 return Err(vm.new_runtime_error(msg.to_owned()));
608 }
609 if self.ag.inner.closed() {
610 self.state.store(AwaitableState::Closed);
611 return Err(vm.new_stop_iteration(None));
612 }
613 self.ag.running_async.store(true, Ordering::Relaxed);
614 self.state.store(AwaitableState::Iter);
615 }
616 AwaitableState::Iter => {}
617 }
618
619 warn_deprecated_throw_signature(&exc_val, &exc_tb, vm)?;
620 let ret = self.ag.inner.throw(
621 self.ag.as_object(),
622 exc_type,
623 exc_val.unwrap_or_none(vm),
624 exc_tb.unwrap_or_none(vm),
625 vm,
626 );
627 if self.aclose && self.ignored_close(&ret) {
628 return Err(self.yield_close(vm));
629 }
630 let res = if self.aclose {
631 ret.and_then(|o| o.into_async_pyresult(vm))
632 } else {
633 PyAsyncGenWrappedValue::unbox(&self.ag, ret, vm)
634 };
635 res.map_err(|e| self.check_error(e, vm))
636 }
637
638 #[pymethod]
639 fn close(&self, vm: &VirtualMachine) -> PyResult<()> {
640 if matches!(self.state.load(), AwaitableState::Closed) {
641 return Ok(());
642 }
643 let result = self.throw(
644 vm.ctx.exceptions.generator_exit.to_owned().into(),
645 OptionalArg::Missing,
646 OptionalArg::Missing,
647 vm,
648 );
649 match result {
650 Ok(_) => Err(vm.new_runtime_error("coroutine ignored GeneratorExit")),
651 Err(e)
652 if e.fast_isinstance(vm.ctx.exceptions.stop_iteration)
653 || e.fast_isinstance(vm.ctx.exceptions.stop_async_iteration)
654 || e.fast_isinstance(vm.ctx.exceptions.generator_exit) =>
655 {
656 Ok(())
657 }
658 Err(e) => Err(e),
659 }
660 }
661}
662
663impl SelfIter for PyAsyncGenAThrow {}
664impl IterNext for PyAsyncGenAThrow {
665 fn next(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
666 PyIterReturn::from_pyresult(zelf.send(vm.ctx.none(), vm), vm)
667 }
668}
669
670impl Destructor for PyAsyncGenAThrow {
671 fn del(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<()> {
672 if matches!(zelf.state.load(), AwaitableState::Init) {
673 let method = if zelf.aclose { "aclose" } else { "athrow" };
674 warn_unawaited_asyncgen_method(&zelf.ag, method, vm);
675 }
676 Ok(())
677 }
678}
679
680#[pyclass(module = false, name = "anext_awaitable", traverse = "manual")]
683#[derive(Debug)]
684pub(crate) struct PyAnextAwaitable {
685 wrapped: PyObjectRef,
686 default_value: PyObjectRef,
687 state: AtomicCell<AwaitableState>,
688}
689
690unsafe impl Traverse for PyAnextAwaitable {
691 fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
692 self.wrapped.traverse(tracer_fn);
693 self.default_value.traverse(tracer_fn);
694 }
695}
696
697impl PyPayload for PyAnextAwaitable {
698 #[inline]
699 fn class(ctx: &Context) -> &'static Py<PyType> {
700 ctx.types.anext_awaitable
701 }
702}
703
704impl PyAnextAwaitable {
705 pub(crate) fn new(wrapped: PyObjectRef, default_value: PyObjectRef) -> Self {
706 Self {
707 wrapped,
708 default_value,
709 state: AtomicCell::new(AwaitableState::Init),
710 }
711 }
712
713 fn check_closed(&self, vm: &VirtualMachine) -> PyResult<()> {
714 if let AwaitableState::Closed = self.state.load() {
715 return Err(vm.new_runtime_error("cannot reuse already awaited __anext__()/asend()"));
716 }
717 Ok(())
718 }
719
720 fn get_awaitable_iter(&self, vm: &VirtualMachine) -> PyResult {
723 use crate::builtins::PyCoroutine;
724 use crate::protocol::PyIter;
725
726 let wrapped = &self.wrapped;
727
728 if wrapped.class().is(vm.ctx.types.async_generator_asend)
730 || wrapped.class().is(vm.ctx.types.async_generator_athrow)
731 {
732 return Ok(wrapped.clone());
733 }
734
735 let awaitable = if wrapped.class().is(vm.ctx.types.coroutine_type) {
737 wrapped.clone()
739 } else {
740 if let Some(generator) = wrapped.downcast_ref::<PyGenerator>()
742 && generator
743 .as_coro()
744 .code()
745 .flags
746 .contains(crate::bytecode::CodeFlags::ITERABLE_COROUTINE)
747 {
748 return Ok(wrapped.clone());
750 }
751 if let Some(await_method) = vm.get_method(wrapped.clone(), identifier!(vm, __await__)) {
753 await_method?.call((), vm)?
754 } else {
755 return Err(vm.new_type_error(format!(
756 "'{}' object can't be awaited",
757 wrapped.class().name()
758 )));
759 }
760 };
761
762 if awaitable.class().is(vm.ctx.types.coroutine_type) {
764 let coro_await = vm.call_method(&awaitable, "__await__", ())?;
765 if !PyIter::check(&coro_await) {
767 return Err(vm.new_type_error("__await__ returned a non-iterable"));
768 }
769 return Ok(coro_await);
770 }
771
772 if awaitable.downcast_ref::<PyCoroutine>().is_some() {
774 return Err(vm.new_type_error("__await__() returned a coroutine"));
775 }
776
777 if !PyIter::check(&awaitable) {
779 return Err(vm.new_type_error(format!(
780 "__await__() returned non-iterator of type '{}'",
781 awaitable.class().name()
782 )));
783 }
784
785 Ok(awaitable)
786 }
787
788 fn handle_result(&self, result: PyResult, vm: &VirtualMachine) -> PyResult {
790 match result {
791 Ok(value) => Ok(value),
792 Err(exc) if exc.fast_isinstance(vm.ctx.exceptions.stop_async_iteration) => {
793 Err(vm.new_stop_iteration(Some(self.default_value.clone())))
794 }
795 Err(exc) => Err(exc),
796 }
797 }
798}
799
800#[pyclass(with(IterNext, Iterable))]
801impl Py<PyAnextAwaitable> {
802 #[pymethod(name = "__await__")]
803 fn r#await(zelf: PyRef<PyAnextAwaitable>, _vm: &VirtualMachine) -> PyRef<PyAnextAwaitable> {
804 zelf
805 }
806
807 #[pymethod]
808 fn send(&self, val: PyObjectRef, vm: &VirtualMachine) -> PyResult {
809 self.check_closed(vm)?;
810 self.state.store(AwaitableState::Iter);
811 let awaitable = self.get_awaitable_iter(vm)?;
812 let result = vm.call_method(&awaitable, "send", (val,));
813 self.handle_result(result, vm)
814 }
815
816 #[pymethod]
817 fn throw(
818 &self,
819 exc_type: PyObjectRef,
820 exc_val: OptionalArg,
821 exc_tb: OptionalArg,
822 vm: &VirtualMachine,
823 ) -> PyResult {
824 self.check_closed(vm)?;
825 warn_deprecated_throw_signature(&exc_val, &exc_tb, vm)?;
826 self.state.store(AwaitableState::Iter);
827 let awaitable = self.get_awaitable_iter(vm)?;
828 let result = vm.call_method(
829 &awaitable,
830 "throw",
831 (
832 exc_type,
833 exc_val.unwrap_or_none(vm),
834 exc_tb.unwrap_or_none(vm),
835 ),
836 );
837 self.handle_result(result, vm)
838 }
839
840 #[pymethod]
841 fn close(&self, vm: &VirtualMachine) {
842 self.state.store(AwaitableState::Closed);
843 if let Ok(awaitable) = self.get_awaitable_iter(vm) {
844 let _ = vm.call_method(&awaitable, "close", ());
845 }
846 }
847}
848
849impl SelfIter for PyAnextAwaitable {}
850impl IterNext for PyAnextAwaitable {
851 fn next(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
852 PyIterReturn::from_pyresult(zelf.send(vm.ctx.none(), vm), vm)
853 }
854}
855
856impl Destructor for PyAsyncGen {
858 fn del(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<()> {
859 if zelf.inner.closed() {
860 return Ok(());
861 }
862 if zelf.ag_finalizer.lock().clone().is_some() && !zelf.ag_closed.load() {
863 Self::call_finalizer(zelf, vm);
864 return Ok(());
865 }
866 if let Err(e) = zelf.inner.close(zelf.as_object(), vm) {
867 crate::coroutine::unraisable_while_closing(zelf.as_object(), &zelf.inner, e, vm);
868 }
869 Ok(())
870 }
871}
872
873impl Drop for PyAsyncGen {
874 fn drop(&mut self) {
875 if let Some(frame) = self.inner.frame_opt() {
876 frame.clear_generator();
877 }
878 }
879}
880
881fn asyncgen_del_needed(zelf: &PyObject) -> bool {
887 let zelf: &Py<PyAsyncGen> = zelf
888 .downcast_ref()
889 .expect("del_needed is only installed on the async_generator type");
890 !zelf.inner.closed.load() && zelf.ag_finalizer.lock().is_some()
891}
892
893pub(crate) fn init(ctx: &'static Context) {
894 PyAsyncGen::extend_class(ctx, ctx.types.async_generator);
895 PyAsyncGenASend::extend_class(ctx, ctx.types.async_generator_asend);
896 PyAsyncGenAThrow::extend_class(ctx, ctx.types.async_generator_athrow);
897 PyAnextAwaitable::extend_class(ctx, ctx.types.anext_awaitable);
898 PyAsyncGenWrappedValue::extend_class(ctx, ctx.types.async_generator_wrapped_value);
899 ctx.types
900 .async_generator
901 .slots
902 .del_needed
903 .store(Some(asyncgen_del_needed));
904}