Skip to main content

rustpython_vm/builtins/
asyncgenerator.rs

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    // whether hooks have been initialized
33    ag_hooks_inited: AtomicCell<bool>,
34    // Distinct from the frame being finished: aclose() sets this
35    // before throwing GeneratorExit, and unwrap sets it on
36    // StopAsyncIteration / GeneratorExit.
37    ag_closed: AtomicCell<bool>,
38    // ag_origin_or_finalizer - stores the finalizer callback
39    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    // Tracked in `make_generator_or_coro`, together with the frame the object
52    // is born owning, so the pair costs one trip through the GC's gen0 list
53    // instead of two.
54    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    /// Initialize async generator hooks.
79    /// Returns Ok(()) if successful, Err if firstiter hook raised an exception.
80    fn init_hooks(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<()> {
81        // = async_gen_init_hooks
82        if zelf.ag_hooks_inited.load() {
83            return Ok(());
84        }
85
86        zelf.ag_hooks_inited.store(true);
87
88        // Get and store finalizer from VM
89        let finalizer = vm.async_gen_finalizer.borrow().clone();
90        if let Some(finalizer) = finalizer {
91            *zelf.ag_finalizer.lock() = Some(finalizer);
92        }
93
94        // Call firstiter hook
95        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    /// Call finalizer hook if set.
105    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            // Create a strong reference for the finalizer call.
111            // This keeps the object alive during the finalizer execution.
112            let obj: PyObjectRef = zelf.to_owned().into();
113
114            // Call the finalizer. Any exceptions are handled as unraisable.
115            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, // already running, all good
357            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// Awaitable wrapper for anext() builtin with default value.
681// When StopAsyncIteration is raised, it converts it to StopIteration(default).
682#[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    /// Get the awaitable iterator from wrapped object.
721    // = anextawaitable_getiter.
722    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 is already an async_generator_asend, it's an iterator
729        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        // _PyCoro_GetAwaitableIter equivalent
736        let awaitable = if wrapped.class().is(vm.ctx.types.coroutine_type) {
737            // Coroutine - get __await__ later
738            wrapped.clone()
739        } else {
740            // Check for generator with CO_ITERABLE_COROUTINE flag
741            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 the generator itself as the iterator
749                return Ok(wrapped.clone());
750            }
751            // Try to get __await__ method
752            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 is a coroutine, get its __await__
763        if awaitable.class().is(vm.ctx.types.coroutine_type) {
764            let coro_await = vm.call_method(&awaitable, "__await__", ())?;
765            // Check that __await__ returned an iterator
766            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        // Check the result is an iterator, not a coroutine
773        if awaitable.downcast_ref::<PyCoroutine>().is_some() {
774            return Err(vm.new_type_error("__await__() returned a coroutine"));
775        }
776
777        // Check that the result is an iterator
778        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    /// Convert StopAsyncIteration to StopIteration(default_value)
789    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
856/// _PyGen_Finalize for async generators
857impl 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
881/// Fast, VM-free check mirroring `<PyAsyncGen as Destructor>::del` (which in
882/// turn mirrors `call_finalizer`): closed, or no finalizer hook installed
883/// (the common case — `sys.set_asyncgen_hooks` is rarely used), means `del`
884/// is a no-op. Skipping the call avoids attaching to a VM (`with_vm`) on
885/// every async generator drop.
886fn 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}