use std::{collections::VecDeque, mem};
use ahash::AHashMap;
use crate::{
asyncio::{Awaiter, CallId, ExternalFutureState, TaskId},
exception_private::RunError,
heap::{ContainsHeap, DropWithHeap, Heap, HeapId, HeapReadOutput, HeapReader},
intern::FunctionId,
resource::{ResourceError, ResourceTracker},
value::Value,
};
pub(crate) const SCHEDULER_TASK_OVERHEAD: usize =
mem::size_of::<(TaskId, Task)>() + mem::size_of::<(HeapId, TaskId)>() + mem::size_of::<TaskId>();
#[derive(Debug, serde::Serialize, serde::Deserialize)]
pub(crate) enum TaskState {
Ready,
Blocked(HeapId),
Completed(Value),
Failed(RunError),
}
impl DropWithHeap for TaskState {
fn drop_with_heap<H: ContainsHeap>(self, heap: &mut H) {
match self {
Self::Ready | Self::Failed(_) => {}
Self::Blocked(id) => heap.heap_mut().dec_ref(id),
Self::Completed(value) => value.drop_with_heap(heap),
}
}
}
#[derive(Debug, serde::Serialize, serde::Deserialize)]
pub(crate) struct Task {
pub id: TaskId,
pub frames: Vec<SerializedTaskFrame>,
pub stack: Vec<Value>,
pub exception_stack: Vec<Value>,
pub instruction_ip: usize,
pub coroutine_id: Option<HeapId>,
pub gather_id: Option<HeapId>,
pub state: TaskState,
}
impl DropWithHeap for Task {
fn drop_with_heap<H: ContainsHeap>(mut self, heap: &mut H) {
for value in self.stack.drain(..) {
value.drop_with_heap(heap);
}
for value in self.exception_stack.drain(..) {
value.drop_with_heap(heap);
}
self.state.drop_with_heap(heap);
if let Some(coro_id) = self.coroutine_id.take() {
heap.heap_mut().dec_ref(coro_id);
}
if let Some(gid) = self.gather_id.take() {
heap.heap_mut().dec_ref(gid);
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub(crate) struct SerializedTaskFrame {
pub function_id: Option<FunctionId>,
pub ip: usize,
pub stack_base: usize,
pub locals_count: u16,
pub exception_stack_base: usize,
pub call_offset: Option<u32>,
#[serde(default)]
pub is_initializer: bool,
}
impl Task {
pub fn new(id: TaskId, coroutine_id: Option<HeapId>, gather_id: Option<HeapId>) -> Self {
Self {
id,
frames: Vec::new(),
stack: Vec::new(),
exception_stack: Vec::new(),
instruction_ip: 0,
coroutine_id,
gather_id,
state: TaskState::Ready,
}
}
#[inline]
pub fn is_finished(&self) -> bool {
matches!(self.state, TaskState::Completed(_) | TaskState::Failed(_))
}
pub(crate) fn saved_context_size(&self) -> usize {
mem::size_of_val(self.frames.as_slice())
+ mem::size_of_val(self.stack.as_slice())
+ mem::size_of_val(self.exception_stack.as_slice())
}
}
#[derive(Debug, serde::Serialize, serde::Deserialize)]
pub(crate) struct Scheduler {
tasks: AHashMap<TaskId, Task>,
ready_queue: VecDeque<TaskId>,
current_task: Option<TaskId>,
next_task_id: u32,
next_call_id: u32,
pending_externals: AHashMap<CallId, HeapId>,
coroutine_to_task: AHashMap<HeapId, TaskId>,
}
impl Scheduler {
pub fn new() -> Self {
let main_task_id = TaskId::default();
let mut main_task = Task::new(main_task_id, None, None);
main_task.state = TaskState::Ready; let mut tasks = AHashMap::new();
tasks.insert(main_task_id, main_task);
Self {
tasks,
ready_queue: VecDeque::new(), current_task: Some(main_task_id),
next_task_id: 1,
next_call_id: 0,
pending_externals: AHashMap::new(),
coroutine_to_task: AHashMap::new(),
}
}
#[inline]
pub fn current_task_id(&self) -> Option<TaskId> {
self.current_task
}
#[inline]
pub fn get_task(&self, task_id: TaskId) -> &Task {
self.tasks.get(&task_id).expect("Scheduler::get_task: task not found")
}
#[inline]
pub fn get_task_mut(&mut self, task_id: TaskId) -> &mut Task {
self.tasks
.get_mut(&task_id)
.expect("Scheduler::get_task_mut: task not found")
}
pub fn allocate_call_id(&mut self) -> CallId {
let id = CallId::new(self.next_call_id);
self.next_call_id += 1;
id
}
pub fn add_pending_external(&mut self, call_id: CallId, future_id: HeapId, heap: &Heap<impl ResourceTracker>) {
heap.inc_ref(future_id);
let prev = self.pending_externals.insert(call_id, future_id);
debug_assert!(prev.is_none(), "add_pending_external: CallId already registered");
}
pub fn take_pending_external(&mut self, call_id: CallId) -> Option<HeapId> {
self.pending_externals.remove(&call_id)
}
pub fn block_current_on(&mut self, awaitable_id: HeapId, heap: &Heap<impl ResourceTracker>) {
if let Some(task_id) = self.current_task {
let task = self.get_task_mut(task_id);
heap.inc_ref(awaitable_id);
task.state = TaskState::Blocked(awaitable_id);
}
}
pub fn pending_call_ids(&self) -> Vec<CallId> {
self.pending_externals.keys().copied().collect()
}
pub fn remove_from_ready_queue(&mut self, task_id: TaskId) {
self.ready_queue.retain(|&id| id != task_id);
}
pub fn spawn(
&mut self,
heap: &Heap<impl ResourceTracker>,
coroutine_id: HeapId,
gather_id: Option<HeapId>,
) -> Result<Option<TaskId>, ResourceError> {
if self.coroutine_to_task.contains_key(&coroutine_id) {
return Ok(None);
}
heap.track_growth(SCHEDULER_TASK_OVERHEAD)?;
let task_id = TaskId::new(self.next_task_id);
self.next_task_id += 1;
heap.inc_ref(coroutine_id);
if let Some(gid) = gather_id {
heap.inc_ref(gid);
}
let task = Task::new(task_id, Some(coroutine_id), gather_id);
self.tasks.insert(task_id, task);
self.coroutine_to_task.insert(coroutine_id, task_id);
self.ready_queue.push_back(task_id);
Ok(Some(task_id))
}
#[inline]
pub fn task_for_coroutine(&self, coroutine_id: HeapId) -> Option<TaskId> {
self.coroutine_to_task.get(&coroutine_id).copied()
}
pub fn next_ready_task(&mut self) -> Option<TaskId> {
self.ready_queue.pop_front()
}
pub fn requeue_ready_front(&mut self, task_id: TaskId) {
self.ready_queue.push_front(task_id);
}
pub fn set_state(&mut self, task_id: TaskId, new_state: TaskState, heap: &mut Heap<impl ResourceTracker>) {
let task = self.get_task_mut(task_id);
let old_state = mem::replace(&mut task.state, new_state);
old_state.drop_with_heap(heap);
}
pub fn make_ready(&mut self, task_id: TaskId, heap: &mut Heap<impl ResourceTracker>) {
self.set_state(task_id, TaskState::Ready, heap);
self.ready_queue.push_back(task_id);
}
pub fn set_current_task(&mut self, task_id: Option<TaskId>) {
self.current_task = task_id;
}
pub fn fail_task(
&mut self,
task_id: TaskId,
error: RunError,
heap: &mut Heap<impl ResourceTracker>,
) -> Option<HeapId> {
let gather_id = self.get_task(task_id).gather_id;
self.set_state(task_id, TaskState::Failed(error), heap);
gather_id
}
pub fn cancel_task(&mut self, task_id: TaskId, heap: &mut HeapReader<'_, impl ResourceTracker>) {
let Some(task) = self.tasks.remove(&task_id) else {
return;
};
if task_id != TaskId::default() {
heap.heap_mut().track_shrink(SCHEDULER_TASK_OVERHEAD);
}
heap.heap_mut().track_shrink(task.saved_context_size());
if self.current_task == Some(task_id) {
self.current_task = None;
}
if let Some(coroutine_id) = task.coroutine_id {
self.coroutine_to_task.remove(&coroutine_id);
}
if !task.is_finished() {
self.ready_queue.retain(|&id| id != task_id);
if let TaskState::Blocked(blocked_id) = task.state
&& let HeapReadOutput::GatherFuture(gather) = heap.read(blocked_id)
{
let inner_task_ids: Vec<TaskId> = gather
.get(heap)
.as_awaited()
.map(|awaited| {
awaited
.pending_children
.keys()
.filter_map(|id| self.coroutine_to_task.get(id).copied())
.collect()
})
.unwrap_or_default();
drop(gather);
for inner_task_id in inner_task_ids {
self.cancel_task(inner_task_id, heap);
}
}
}
task.drop_with_heap(heap);
}
#[must_use]
pub fn fail_for_call(
&mut self,
call_id: CallId,
error: &RunError,
heap: &mut HeapReader<'_, impl ResourceTracker>,
) -> Option<Awaiter> {
let future_id = self.pending_externals.remove(&call_id)?;
let HeapReadOutput::ExternalFuture(mut fut) = heap.read(future_id) else {
panic!("pending_externals entry doesn't point to an ExternalFuture")
};
let awaiter = match mem::replace(&mut fut.get_mut(heap).state, ExternalFutureState::Failed(error.clone())) {
ExternalFutureState::Pending { awaiter } => awaiter,
ExternalFutureState::Resolved(_) | ExternalFutureState::Failed(_) => {
panic!("fail_for_call: future was already resolved")
}
};
drop(fut);
heap.dec_ref(future_id);
match awaiter {
None => None,
Some(Awaiter::Task(task_id)) => {
let gather_id = self.tasks.get(&task_id).and_then(|t| t.gather_id);
if let Some(gather_id) = gather_id {
let HeapReadOutput::GatherFuture(mut gather_rd) = heap.read(gather_id) else {
panic!("gather_id doesn't point to a GatherFuture")
};
let outer_awaiter = gather_rd.fail(self, heap, error);
drop(gather_rd);
Some(outer_awaiter)
} else {
Some(Awaiter::Task(task_id))
}
}
Some(Awaiter::GatherSlot { gather, .. }) => {
let HeapReadOutput::GatherFuture(mut gather_rd) = heap.read(gather) else {
panic!("gather_id doesn't point to a GatherFuture")
};
let outer_awaiter = gather_rd.fail(self, heap, error);
drop(gather_rd);
heap.dec_ref(gather);
Some(outer_awaiter)
}
}
}
#[inline]
pub fn is_task_failed(&self, task_id: TaskId) -> bool {
self.tasks
.get(&task_id)
.is_some_and(|task| matches!(task.state, TaskState::Failed(_)))
}
#[inline]
pub fn has_task(&self, task_id: TaskId) -> bool {
self.tasks.contains_key(&task_id)
}
pub fn cleanup(&mut self, heap: &mut HeapReader<'_, impl ResourceTracker>) {
for (_, future_id) in mem::take(&mut self.pending_externals) {
heap.dec_ref(future_id);
}
let task_ids: Vec<TaskId> = self.tasks.keys().copied().collect();
for task_id in task_ids {
self.cancel_task(task_id, heap);
}
}
}
impl Default for Scheduler {
fn default() -> Self {
Self::new()
}
}