use crate::application::submit_to_main_thread;
use crate::sys;
use std::cell::Cell;
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::task::{Context, RawWaker, RawWakerVTable};
static NEXT_TASK_ID: AtomicUsize = AtomicUsize::new(1);
struct Inner {
task_id: usize,
}
impl Inner {
fn new(task_id: usize) -> Self {
Inner { task_id }
}
}
struct Waker {
inner: Arc<Inner>,
}
const WAKER_VTABLE: RawWakerVTable = RawWakerVTable::new(
|data| {
let w = unsafe { Arc::from_raw(data as *const Waker) };
let w2 = w.clone();
_ = Arc::into_raw(w); RawWaker::new(Arc::into_raw(w2) as *const (), &WAKER_VTABLE)
},
|data| {
let w = unsafe { Arc::from_raw(data as *const Waker) };
wake_task(w.inner.task_id);
},
|data| {
let w = unsafe { Arc::from_raw(data as *const Waker) };
wake_task(w.inner.task_id);
std::mem::forget(w);
},
|data| {
let w = unsafe { Arc::from_raw(data as *const Waker) };
drop(w);
},
);
impl Waker {
fn into_waker(self) -> std::task::Waker {
let arc_waker = Arc::into_raw(Arc::new(self));
unsafe { std::task::Waker::from_raw(RawWaker::new(arc_waker as *const (), &WAKER_VTABLE)) }
}
}
struct Task {
context: logwise::ContextToken,
our_task_id: usize,
future: Pin<Box<dyn Future<Output = ()> + 'static>>,
wake_inner: Arc<Inner>,
}
fn wake_task(task_id: usize) {
crate::application::submit_to_main_thread("wake_task".to_string(), move || {
let mut pollable = POLLABLE.take();
pollable.push(task_id);
POLLABLE.replace(pollable);
main_executor_iter();
});
}
thread_local! {
static RUNNING: Cell<Option<HashMap<usize, Task>>> = const { Cell::new(None) };
static POLLABLE: Cell<Vec<usize>> = const { Cell::new(Vec::new()) };
static IN_FLIGHT: Cell<Vec<usize>> = const { Cell::new(Vec::new()) };
}
pub async fn on_main_thread_async<R: Send + 'static, F: Future<Output = R> + Send + 'static>(
debug_label: String,
future: F,
) -> R {
let (sender, fut) = r#continue::continuation();
crate::application::submit_to_main_thread(debug_label.clone(), || {
already_on_main_thread_submit(debug_label, async move {
let r = future.await;
sender.send(r);
})
});
fut.await
}
pub fn already_on_main_thread_submit<F: Future<Output = ()> + 'static>(
debug_label: String,
future: F,
) {
assert!(sys::is_main_thread());
let task_id = NEXT_TASK_ID.fetch_add(1, Ordering::Relaxed);
let wake_inner = Arc::new(Inner::new(task_id));
let new_context = logwise::context::child(logwise::context::capture(), "app_window.task");
logwise::log!(
"app_window: creating task {id} {label}",
id = task_id,
label = debug_label
);
let task = Task {
our_task_id: task_id,
context: new_context,
future: Box::pin(future),
wake_inner,
};
let mut pollable = POLLABLE.take();
pollable.push(task_id);
POLLABLE.replace(pollable);
let mut running = RUNNING.take().unwrap_or_default();
running.insert(task_id, task);
RUNNING.replace(Some(running));
main_executor_iter();
}
fn main_executor_iter() {
let begin_iter = crate::application::time::Instant::now();
let mut swap_pollable = POLLABLE.take();
let poll = if swap_pollable.is_empty() {
None
} else {
Some(swap_pollable.remove(0))
};
POLLABLE.replace(swap_pollable);
match poll {
None => {
}
Some(woke_task_id) => {
let in_flight = IN_FLIGHT.take();
let is_in_flight = in_flight.contains(&woke_task_id);
IN_FLIGHT.replace(in_flight);
if is_in_flight {
let mut pollable = POLLABLE.take();
pollable.push(woke_task_id);
POLLABLE.replace(pollable);
return;
}
let mut running = RUNNING.take().unwrap_or_default();
let task = running.remove(&woke_task_id);
RUNNING.replace(Some(running));
let mut task = match task {
Some(task) => task,
None => {
return;
}
};
let task_id = task.our_task_id;
let mut in_flight = IN_FLIGHT.take();
in_flight.push(task.our_task_id);
IN_FLIGHT.replace(in_flight);
let waker = Waker {
inner: task.wake_inner.clone(),
};
let into_waker = waker.into_waker();
let mut context = Context::from_waker(&into_waker);
let poll_result = {
let _entered = logwise::context::enter(task.context);
task.future.as_mut().poll(&mut context)
};
let mut in_flight = IN_FLIGHT.take();
in_flight.retain(|&id| id != task.our_task_id);
IN_FLIGHT.replace(in_flight);
match poll_result {
std::task::Poll::Ready(()) => {
}
std::task::Poll::Pending => {
let mut running = RUNNING.take().unwrap_or_default();
running.insert(task.our_task_id, task);
RUNNING.replace(Some(running));
}
}
submit_to_main_thread("main_executor_iter".to_string(), main_executor_iter);
if begin_iter.elapsed() > crate::application::time::Duration::from_millis(10) {
logwise::event!(
class: performance,
severity: warn,
name: "app_window.main_executor.poll_overran",
task = support(task_id as u64),
duration_ms = support(begin_iter.elapsed().as_millis() as u64),
);
}
}
}
}