use std::collections::HashMap;
use std::sync::Arc;
use shape_value::heap_value::{HeapKind, TaskGroupData};
use shape_value::{NativeKind, VMError};
use crate::executor::vm_impl::stack::{clone_with_kind, drop_with_kind};
type Kinded = (u64, NativeKind);
#[derive(Debug, Clone)]
pub enum TaskStatus {
Pending,
Completed(Kinded),
Cancelled,
}
pub struct TaskScheduler {
callables: HashMap<u64, Kinded>,
results: HashMap<u64, TaskStatus>,
external_receivers: HashMap<u64, tokio::sync::oneshot::Receiver<Result<Kinded, String>>>,
}
impl TaskScheduler {
pub fn new() -> Self {
Self {
callables: HashMap::new(),
results: HashMap::new(),
external_receivers: HashMap::new(),
}
}
pub fn register(&mut self, task_id: u64, callable_bits: u64, callable_kind: NativeKind) {
if let Some((old_bits, old_kind)) = self.callables.remove(&task_id) {
drop_with_kind(old_bits, old_kind);
}
self.callables
.insert(task_id, (callable_bits, callable_kind));
self.results.insert(task_id, TaskStatus::Pending);
}
pub fn take_callable(&mut self, task_id: u64) -> Option<Kinded> {
self.callables.remove(&task_id)
}
pub fn complete(&mut self, task_id: u64, value_bits: u64, value_kind: NativeKind) {
if let Some(TaskStatus::Completed((old_bits, old_kind))) =
self.results.insert(task_id, TaskStatus::Completed((value_bits, value_kind)))
{
drop_with_kind(old_bits, old_kind);
}
}
pub fn cancel(&mut self, task_id: u64) {
if let Some(TaskStatus::Pending) = self.results.get(&task_id) {
self.results.insert(task_id, TaskStatus::Cancelled);
if let Some((bits, kind)) = self.callables.remove(&task_id) {
drop_with_kind(bits, kind);
}
}
}
pub fn get_result(&self, task_id: u64) -> Option<&TaskStatus> {
self.results.get(&task_id)
}
pub fn is_resolved(&self, task_id: u64) -> bool {
matches!(
self.results.get(&task_id),
Some(TaskStatus::Completed(_)) | Some(TaskStatus::Cancelled)
)
}
pub fn register_external(
&mut self,
task_id: u64,
) -> tokio::sync::oneshot::Sender<Result<Kinded, String>> {
let (tx, rx) = tokio::sync::oneshot::channel();
self.results.insert(task_id, TaskStatus::Pending);
self.external_receivers.insert(task_id, rx);
tx
}
pub fn try_resolve_external(&mut self, task_id: u64) -> Option<Result<Kinded, VMError>> {
if let Some(TaskStatus::Completed((bits, kind))) = self.results.get(&task_id).cloned() {
clone_with_kind(bits, kind);
return Some(Ok((bits, kind)));
}
if let Some(rx) = self.external_receivers.get_mut(&task_id) {
match rx.try_recv() {
Ok(Ok((bits, kind))) => {
clone_with_kind(bits, kind);
self.results
.insert(task_id, TaskStatus::Completed((bits, kind)));
self.external_receivers.remove(&task_id);
Some(Ok((bits, kind)))
}
Ok(Err(e)) => {
self.external_receivers.remove(&task_id);
Some(Err(VMError::RuntimeError(e)))
}
Err(tokio::sync::oneshot::error::TryRecvError::Empty) => None,
Err(tokio::sync::oneshot::error::TryRecvError::Closed) => {
self.external_receivers.remove(&task_id);
Some(Err(VMError::RuntimeError(
"Remote task cancelled".to_string(),
)))
}
}
} else {
None
}
}
pub fn has_external(&self, task_id: u64) -> bool {
self.external_receivers.contains_key(&task_id)
}
pub fn take_external_receiver(
&mut self,
task_id: u64,
) -> Option<tokio::sync::oneshot::Receiver<Result<Kinded, String>>> {
self.external_receivers.remove(&task_id)
}
pub fn resolve_task<F>(&mut self, task_id: u64, executor_fn: F) -> Result<Kinded, VMError>
where
F: FnOnce(Kinded) -> Result<Kinded, VMError>,
{
if let Some(TaskStatus::Completed((bits, kind))) = self.results.get(&task_id).cloned() {
clone_with_kind(bits, kind);
return Ok((bits, kind));
}
if let Some(TaskStatus::Cancelled) = self.results.get(&task_id) {
return Err(VMError::RuntimeError(format!(
"Task {} was cancelled",
task_id
)));
}
let callable = self.take_callable(task_id).ok_or_else(|| {
VMError::RuntimeError(format!("No callable registered for task {}", task_id))
})?;
let (bits, kind) = executor_fn(callable)?;
clone_with_kind(bits, kind);
self.results
.insert(task_id, TaskStatus::Completed((bits, kind)));
Ok((bits, kind))
}
pub fn resolve_task_group<F>(
&mut self,
kind: u8,
task_ids: &[u64],
mut executor_fn: F,
) -> Result<Kinded, VMError>
where
F: FnMut(Kinded) -> Result<Kinded, VMError>,
{
match kind {
0 => {
for &id in task_ids {
let (bits, k) = self.resolve_task(id, &mut executor_fn)?;
drop_with_kind(bits, k);
}
let bits = Arc::into_raw(Arc::new(TaskGroupData {
kind: 0,
task_ids: task_ids.to_vec(),
})) as u64;
Ok((bits, NativeKind::Ptr(HeapKind::TaskGroup)))
}
1 => {
for &id in task_ids {
let res = self.resolve_task(id, &mut executor_fn)?;
return Ok(res);
}
Err(VMError::RuntimeError(
"Race join with empty task list".to_string(),
))
}
2 => {
let mut last_err = None;
for &id in task_ids {
match self.resolve_task(id, &mut executor_fn) {
Ok(res) => return Ok(res),
Err(e) => last_err = Some(e),
}
}
Err(last_err.unwrap_or_else(|| {
VMError::RuntimeError("Any join with empty task list".to_string())
}))
}
3 => {
for &id in task_ids {
if let Ok((bits, k)) = self.resolve_task(id, &mut executor_fn) {
drop_with_kind(bits, k);
}
}
let bits = Arc::into_raw(Arc::new(TaskGroupData {
kind: 3,
task_ids: task_ids.to_vec(),
})) as u64;
Ok((bits, NativeKind::Ptr(HeapKind::TaskGroup)))
}
_ => Err(VMError::RuntimeError(format!(
"Unknown join kind: {}",
kind
))),
}
}
}
#[cfg(feature = "gc")]
impl TaskScheduler {
pub(crate) fn scan_roots(&self, _visitor: &mut dyn FnMut(*mut u8)) {
todo!(
"phase-2c — ADR-006 §2.7.4: kinded GC root walker for TaskScheduler. \
The pre-bulldozer trace_nanboxed_bits path decoded ValueWord tag \
bits; the kinded equivalent (parallel kinds track + per-HeapKind \
dispatch via slot.as_heap_value()) belongs to the Phase-2c GC \
rebuild and is out of R-async-time scope."
)
}
}
impl Drop for TaskScheduler {
fn drop(&mut self) {
for (_, (bits, kind)) in self.callables.drain() {
drop_with_kind(bits, kind);
}
for (_, status) in self.results.drain() {
if let TaskStatus::Completed((bits, kind)) = status {
drop_with_kind(bits, kind);
}
}
}
}
impl std::fmt::Debug for TaskScheduler {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TaskScheduler")
.field("callables", &format!("[{} pending]", self.callables.len()))
.field("results", &format!("[{} entries]", self.results.len()))
.field(
"external_receivers",
&format!("[{} pending]", self.external_receivers.len()),
)
.finish()
}
}
impl Default for TaskScheduler {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn function_callable(func_id: u64) -> Kinded {
(func_id, NativeKind::Ptr(HeapKind::Future))
}
fn float_result(v: f64) -> Kinded {
(v.to_bits(), NativeKind::Float64)
}
#[test]
fn test_register_and_take_callable() {
let mut sched = TaskScheduler::new();
let (bits, kind) = function_callable(42);
sched.register(1, bits, kind);
assert!(matches!(sched.get_result(1), Some(TaskStatus::Pending)));
let callable = sched.take_callable(1);
assert!(callable.is_some());
assert!(sched.take_callable(1).is_none());
}
#[test]
fn test_resolve_task_synchronous() {
let mut sched = TaskScheduler::new();
let (b, k) = function_callable(0);
sched.register(1, b, k);
let result = sched.resolve_task(1, |_callable| Ok(float_result(99.0)));
assert!(result.is_ok());
let (bits, kind) = result.unwrap();
assert_eq!(kind, NativeKind::Float64);
assert!((f64::from_bits(bits) - 99.0).abs() < f64::EPSILON);
let cached = sched.resolve_task(1, |_| panic!("should not be called"));
assert!(cached.is_ok());
}
#[test]
fn test_cancel_task() {
let mut sched = TaskScheduler::new();
let (b, k) = function_callable(0);
sched.register(1, b, k);
sched.cancel(1);
assert!(sched.is_resolved(1));
let result = sched.resolve_task(1, |_| Ok(float_result(0.0)));
assert!(result.is_err());
}
#[test]
fn test_resolve_all_group() {
let mut sched = TaskScheduler::new();
let (b1, k1) = function_callable(0);
let (b2, k2) = function_callable(1);
sched.register(1, b1, k1);
sched.register(2, b2, k2);
let mut call_count = 0u32;
let result = sched.resolve_task_group(0, &[1, 2], |_callable| {
call_count += 1;
Ok(float_result(call_count as f64))
});
assert!(result.is_ok());
let (_bits, kind) = result.unwrap();
assert_eq!(kind, NativeKind::Ptr(HeapKind::TaskGroup));
assert_eq!(call_count, 2);
}
#[test]
fn test_resolve_race_group() {
let mut sched = TaskScheduler::new();
let (b1, k1) = function_callable(0);
let (b2, k2) = function_callable(1);
sched.register(10, b1, k1);
sched.register(20, b2, k2);
let result = sched.resolve_task_group(1, &[10, 20], |_| Ok(float_result(7.0)));
assert!(result.is_ok());
let (bits, kind) = result.unwrap();
assert_eq!(kind, NativeKind::Float64);
assert!((f64::from_bits(bits) - 7.0).abs() < f64::EPSILON);
}
#[test]
fn test_register_external_and_resolve() {
let mut sched = TaskScheduler::new();
let tx = sched.register_external(100);
assert!(sched.has_external(100));
assert!(matches!(sched.get_result(100), Some(TaskStatus::Pending)));
assert!(sched.try_resolve_external(100).is_none());
tx.send(Ok(float_result(42.0))).unwrap();
let result = sched.try_resolve_external(100);
assert!(result.is_some());
let (bits, kind) = result.unwrap().unwrap();
assert_eq!(kind, NativeKind::Float64);
assert!((f64::from_bits(bits) - 42.0).abs() < f64::EPSILON);
assert!(!sched.has_external(100));
}
#[test]
fn test_external_task_error() {
let mut sched = TaskScheduler::new();
let tx = sched.register_external(200);
tx.send(Err("connection refused".to_string())).unwrap();
let result = sched.try_resolve_external(200);
assert!(result.is_some());
assert!(result.unwrap().is_err());
}
#[test]
fn test_external_task_cancelled() {
let mut sched = TaskScheduler::new();
let tx = sched.register_external(300);
drop(tx);
let result = sched.try_resolve_external(300);
assert!(result.is_some());
assert!(result.unwrap().is_err());
}
#[test]
fn test_take_external_receiver() {
let mut sched = TaskScheduler::new();
let _tx = sched.register_external(400);
assert!(sched.has_external(400));
let rx = sched.take_external_receiver(400);
assert!(rx.is_some());
assert!(!sched.has_external(400));
}
}