use core::future::Future;
use core::pin::Pin;
use core::task::{Context, Poll};
use crate::error::TaskError;
use crate::platform::*;
use crate::{TaskContext, TaskId};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CoroutineState {
Created,
Ready,
Running,
Yielded,
Completed,
Error,
}
pub trait Coroutine {
type Yield;
type Return;
fn resume(&mut self) -> CoroutineResult<Self::Yield, Self::Return>;
fn state(&self) -> CoroutineState;
fn is_resumable(&self) -> bool {
matches!(
self.state(),
CoroutineState::Created | CoroutineState::Ready | CoroutineState::Yielded
)
}
}
pub enum CoroutineResult<Y, R> {
Yielded(Y),
Complete(R),
Error(TaskError),
}
pub struct FunctionCoroutine<Y, R> {
state_fn: Option<Box<dyn FnMut() -> CoroutineResult<Y, R> + Send>>,
state: CoroutineState,
_context: TaskContext,
}
impl<Y, R> FunctionCoroutine<Y, R>
where
Y: Send + 'static,
R: Send + 'static,
{
pub fn new<F>(func: F) -> Self
where
F: FnMut() -> CoroutineResult<Y, R> + Send + 'static,
{
Self {
state_fn: Some(Box::new(func)),
state: CoroutineState::Created,
_context: TaskContext::new(TaskId::new(0)),
}
}
}
impl<Y, R> Coroutine for FunctionCoroutine<Y, R>
where
Y: Send + 'static,
R: Send + 'static,
{
type Yield = Y;
type Return = R;
fn resume(&mut self) -> CoroutineResult<Self::Yield, Self::Return> {
if let Some(mut func) = self.state_fn.take() {
self.state = CoroutineState::Running;
let result = func();
match &result {
CoroutineResult::Yielded(_) => {
self.state = CoroutineState::Yielded;
self.state_fn = Some(func);
}
CoroutineResult::Complete(_) => {
self.state = CoroutineState::Completed;
}
CoroutineResult::Error(_) => {
self.state = CoroutineState::Error;
}
}
result
} else {
CoroutineResult::Error(TaskError::InvalidState)
}
}
fn state(&self) -> CoroutineState {
self.state
}
}
pub struct CoroutineIterator<C> {
coroutine: C,
}
impl<C> CoroutineIterator<C>
where
C: Coroutine,
{
pub fn new(coroutine: C) -> Self {
Self { coroutine }
}
}
impl<C> Iterator for CoroutineIterator<C>
where
C: Coroutine,
{
type Item = C::Yield;
fn next(&mut self) -> Option<Self::Item> {
if !self.coroutine.is_resumable() {
return None;
}
match self.coroutine.resume() {
CoroutineResult::Yielded(value) => Some(value),
CoroutineResult::Complete(_) | CoroutineResult::Error(_) => None,
}
}
}
pub struct CoroutineFuture<C> {
coroutine: Option<C>,
}
impl<C> CoroutineFuture<C>
where
C: Coroutine,
{
pub fn new(coroutine: C) -> Self {
Self {
coroutine: Some(coroutine),
}
}
}
impl<C> Future for CoroutineFuture<C>
where
C: Coroutine + Unpin,
{
type Output = Result<C::Return, TaskError>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Some(coroutine) = self.coroutine.as_mut() else {
return Poll::Ready(Err(TaskError::AlreadyCompleted));
};
match coroutine.resume() {
CoroutineResult::Yielded(_) => {
cx.waker().wake_by_ref();
Poll::Pending
}
CoroutineResult::Complete(value) => {
self.coroutine = None;
Poll::Ready(Ok(value))
}
CoroutineResult::Error(e) => {
self.coroutine = None;
Poll::Ready(Err(e))
}
}
}
}
pub trait CoroutineExt: Sized {
type Yield;
type Return;
fn into_coroutine(self) -> FunctionCoroutine<Self::Yield, Self::Return>;
}
impl<F, Y, R> CoroutineExt for F
where
F: FnMut() -> CoroutineResult<Y, R> + Send + 'static,
Y: Send + 'static,
R: Send + 'static,
{
type Yield = Y;
type Return = R;
fn into_coroutine(self) -> FunctionCoroutine<Y, R> {
FunctionCoroutine::new(self)
}
}
#[macro_export]
macro_rules! coroutine {
($($body:tt)*) => {{
move || {
$($body)*
}
}};
}
#[macro_export]
macro_rules! co_yield {
($value:expr) => {{
return $crate::coroutine::CoroutineResult::Yielded($value);
}};
}
#[macro_export]
macro_rules! co_return {
($value:expr) => {{
return $crate::coroutine::CoroutineResult::Complete($value);
}};
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_coroutine() {
let mut counter = 0;
let mut coro = FunctionCoroutine::new(move || {
counter += 1;
if counter < 3 {
CoroutineResult::Yielded(counter)
} else {
CoroutineResult::Complete(counter)
}
});
assert_eq!(coro.state(), CoroutineState::Created);
match coro.resume() {
CoroutineResult::Yielded(1) => {}
_ => panic!("Expected yield of 1"),
}
match coro.resume() {
CoroutineResult::Yielded(2) => {}
_ => panic!("Expected yield of 2"),
}
match coro.resume() {
CoroutineResult::Complete(3) => {}
_ => panic!("Expected completion with 3"),
}
assert_eq!(coro.state(), CoroutineState::Completed);
}
#[test]
fn test_coroutine_iterator() {
let mut counter = 0;
let coro = FunctionCoroutine::new(move || {
counter += 1;
if counter <= 3 {
CoroutineResult::Yielded(counter)
} else {
CoroutineResult::Complete(())
}
});
let iter = CoroutineIterator::new(coro);
let values: Vec<i32> = iter.collect();
assert_eq!(values, vec![1, 2, 3]);
}
}