use crate::{
bytecode::{Instruction, OpCode, Operand},
executor::VirtualMachine,
executor::vm_impl::stack::drop_with_kind,
};
use shape_value::{
NativeKind, VMError,
heap_value::{HeapKind, TaskGroupData},
};
use std::sync::Arc;
#[derive(Debug, Clone)]
pub enum AsyncExecutionResult {
Continue,
Yielded,
Suspended(SuspensionInfo),
}
#[derive(Debug, Clone)]
pub struct SuspensionInfo {
pub wait_type: WaitType,
pub resume_ip: usize,
}
#[derive(Debug, Clone)]
pub enum WaitType {
NextBar { source: String },
Timer { id: u64 },
AnyEvent,
Future { id: u64 },
TaskGroup { kind: u8, task_ids: Vec<u64> },
}
impl VirtualMachine {
#[inline(always)]
pub(in crate::executor) fn exec_async_op(
&mut self,
instruction: &Instruction,
) -> Result<AsyncExecutionResult, VMError> {
use OpCode::*;
match instruction.opcode {
Yield => self.op_yield(),
Suspend => self.op_suspend(instruction),
Resume => self.op_resume(instruction),
Poll => self.op_poll(),
AwaitBar => self.op_await_bar(instruction),
AwaitTick => self.op_await_tick(instruction),
EmitAlert => self.op_emit_alert(),
EmitEvent => self.op_emit_event(),
Await => self.op_await(),
SpawnTask => self.op_spawn_task(),
JoinInit => self.op_join_init(instruction),
JoinAwait => self.op_join_await(),
CancelTask => self.op_cancel_task(),
AsyncScopeEnter => self.op_async_scope_enter(),
AsyncScopeExit => self.op_async_scope_exit(),
_ => unreachable!(
"exec_async_op called with non-async opcode: {:?}",
instruction.opcode
),
}
}
fn op_yield(&mut self) -> Result<AsyncExecutionResult, VMError> {
Ok(AsyncExecutionResult::Yielded)
}
fn op_suspend(&mut self, instruction: &Instruction) -> Result<AsyncExecutionResult, VMError> {
let wait_type = match &instruction.operand {
Some(Operand::Const(idx)) => {
let _ = idx;
WaitType::AnyEvent
}
_ => WaitType::AnyEvent,
};
Ok(AsyncExecutionResult::Suspended(SuspensionInfo {
wait_type,
resume_ip: self.ip,
}))
}
fn op_resume(&mut self, _instruction: &Instruction) -> Result<AsyncExecutionResult, VMError> {
Ok(AsyncExecutionResult::Continue)
}
fn op_poll(&mut self) -> Result<AsyncExecutionResult, VMError> {
self.push_kinded(0u64, NativeKind::Null)?;
Ok(AsyncExecutionResult::Continue)
}
fn op_await_bar(&mut self, instruction: &Instruction) -> Result<AsyncExecutionResult, VMError> {
let source = match &instruction.operand {
Some(Operand::Const(idx)) => {
match self.program.constants.get(*idx as usize) {
Some(crate::bytecode::Constant::String(s)) => s.clone(),
_ => "default".to_string(),
}
}
_ => "default".to_string(),
};
Ok(AsyncExecutionResult::Suspended(SuspensionInfo {
wait_type: WaitType::NextBar { source },
resume_ip: self.ip,
}))
}
fn op_await_tick(
&mut self,
instruction: &Instruction,
) -> Result<AsyncExecutionResult, VMError> {
let timer_id = match &instruction.operand {
Some(Operand::Const(idx)) => {
match self.program.constants.get(*idx as usize) {
Some(crate::bytecode::Constant::Number(n)) => *n as u64,
_ => 0,
}
}
_ => 0,
};
Ok(AsyncExecutionResult::Suspended(SuspensionInfo {
wait_type: WaitType::Timer { id: timer_id },
resume_ip: self.ip,
}))
}
fn op_emit_alert(&mut self) -> Result<AsyncExecutionResult, VMError> {
let (bits, kind) = self.pop_kinded()?;
drop_with_kind(bits, kind);
Ok(AsyncExecutionResult::Continue)
}
fn op_await(&mut self) -> Result<AsyncExecutionResult, VMError> {
let sp_before = self.sp;
let (bits, kind) = self.pop_kinded()?;
match kind {
NativeKind::Ptr(HeapKind::Future) => {
let task_id = bits;
let result = self.resolve_spawned_task(task_id)?;
self.push_kinded(result.raw(), result.kind())?;
std::mem::forget(result);
debug_assert_eq!(
self.sp, sp_before,
"op_await (Future): stack depth changed (before={}, after={})",
sp_before, self.sp
);
Ok(AsyncExecutionResult::Continue)
}
_ => {
self.push_kinded(bits, kind)?;
debug_assert_eq!(
self.sp, sp_before,
"op_await (sync shortcut): stack depth changed (before={}, after={})",
sp_before, self.sp
);
Ok(AsyncExecutionResult::Continue)
}
}
}
fn op_spawn_task(&mut self) -> Result<AsyncExecutionResult, VMError> {
let sp_before = self.sp;
let (slot_bits, slot_kind) = self.pop_kinded()?;
let task_id = self.next_future_id();
match slot_kind {
NativeKind::Ptr(HeapKind::Closure) | NativeKind::UInt64 => {
self.task_scheduler.register(task_id, slot_bits, slot_kind);
}
_ => {
self.task_scheduler.complete(task_id, slot_bits, slot_kind);
}
}
if let Some(scope) = self.async_scope_stack.last_mut() {
scope.push(task_id);
}
self.push_kinded(task_id, NativeKind::Ptr(HeapKind::Future))?;
debug_assert_eq!(
self.sp, sp_before,
"op_spawn_task: stack depth changed (before={}, after={})",
sp_before, self.sp
);
Ok(AsyncExecutionResult::Continue)
}
fn op_join_init(&mut self, instruction: &Instruction) -> Result<AsyncExecutionResult, VMError> {
let packed = match &instruction.operand {
Some(Operand::Count(n)) => *n,
_ => {
return Err(VMError::RuntimeError(
"JoinInit requires Count operand".to_string(),
));
}
};
let kind = ((packed >> 14) & 0x03) as u8;
let arity = (packed & 0x3FFF) as usize;
if self.sp < arity {
return Err(VMError::StackUnderflow);
}
let mut task_ids: Vec<u64> = Vec::with_capacity(arity);
for _ in 0..arity {
let (bits, slot_kind) = self.pop_kinded()?;
match slot_kind {
NativeKind::Ptr(HeapKind::Future) => {
task_ids.push(bits);
}
_ => {
drop_with_kind(bits, slot_kind);
return Err(VMError::RuntimeError(format!(
"JoinInit expected Future, got {:?}",
slot_kind
)));
}
}
}
task_ids.reverse();
let arc: Arc<TaskGroupData> = Arc::new(TaskGroupData { kind, task_ids });
let bits = Arc::into_raw(arc) as u64;
self.push_kinded(bits, NativeKind::Ptr(HeapKind::TaskGroup))?;
Ok(AsyncExecutionResult::Continue)
}
fn op_join_await(&mut self) -> Result<AsyncExecutionResult, VMError> {
let sp_before = self.sp;
let (bits, slot_kind) = self.pop_kinded()?;
match slot_kind {
NativeKind::Ptr(HeapKind::TaskGroup) => {
let arc: Arc<TaskGroupData> =
unsafe { Arc::from_raw(bits as *const TaskGroupData) };
let join_kind = arc.kind;
let task_ids = arc.task_ids.clone();
drop(arc);
match join_kind {
0 => {
for &id in &task_ids {
let result = self.resolve_spawned_task(id)?;
drop_with_kind(result.raw(), result.kind());
std::mem::forget(result);
}
let aggregate: Arc<TaskGroupData> = Arc::new(TaskGroupData {
kind: 0,
task_ids: task_ids.clone(),
});
let result_bits = Arc::into_raw(aggregate) as u64;
self.push_kinded(
result_bits,
NativeKind::Ptr(HeapKind::TaskGroup),
)?;
}
1 => {
let mut pushed = false;
for (idx, &id) in task_ids.iter().enumerate() {
let result = self.resolve_spawned_task(id)?;
if idx == 0 {
self.push_kinded(result.raw(), result.kind())?;
std::mem::forget(result);
pushed = true;
} else {
drop_with_kind(result.raw(), result.kind());
std::mem::forget(result);
}
}
if !pushed {
return Err(VMError::RuntimeError(
"Race join with empty task list".to_string(),
));
}
}
2 => {
let mut last_err: Option<VMError> = None;
let mut pushed = false;
for &id in &task_ids {
match self.resolve_spawned_task(id) {
Ok(result) => {
self.push_kinded(result.raw(), result.kind())?;
std::mem::forget(result);
pushed = true;
break;
}
Err(e) => last_err = Some(e),
}
}
if !pushed {
return 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(result) = self.resolve_spawned_task(id) {
drop_with_kind(result.raw(), result.kind());
std::mem::forget(result);
}
}
let aggregate: Arc<TaskGroupData> = Arc::new(TaskGroupData {
kind: 3,
task_ids: task_ids.clone(),
});
let result_bits = Arc::into_raw(aggregate) as u64;
self.push_kinded(
result_bits,
NativeKind::Ptr(HeapKind::TaskGroup),
)?;
}
other => {
return Err(VMError::RuntimeError(format!(
"Unknown join kind: {}",
other
)));
}
}
debug_assert_eq!(
self.sp, sp_before,
"op_join_await: stack depth changed (before={}, after={})",
sp_before, self.sp
);
Ok(AsyncExecutionResult::Continue)
}
_ => {
drop_with_kind(bits, slot_kind);
Err(VMError::RuntimeError(format!(
"JoinAwait expected TaskGroup, got {:?}",
slot_kind
)))
}
}
}
fn op_cancel_task(&mut self) -> Result<AsyncExecutionResult, VMError> {
let (bits, slot_kind) = self.pop_kinded()?;
match slot_kind {
NativeKind::Ptr(HeapKind::Future) => {
let id = bits;
self.task_scheduler.cancel(id);
Ok(AsyncExecutionResult::Continue)
}
_ => {
drop_with_kind(bits, slot_kind);
Err(VMError::RuntimeError(format!(
"CancelTask expected Future, got {:?}",
slot_kind
)))
}
}
}
fn op_async_scope_enter(&mut self) -> Result<AsyncExecutionResult, VMError> {
let depth_before = self.async_scope_stack.len();
self.async_scope_stack.push(Vec::new());
debug_assert_eq!(
self.async_scope_stack.len(),
depth_before + 1,
"op_async_scope_enter: scope stack depth not incremented"
);
Ok(AsyncExecutionResult::Continue)
}
fn op_async_scope_exit(&mut self) -> Result<AsyncExecutionResult, VMError> {
debug_assert!(
!self.async_scope_stack.is_empty(),
"op_async_scope_exit: scope stack is empty (mismatched Enter/Exit)"
);
if let Some(mut scope_tasks) = self.async_scope_stack.pop() {
scope_tasks.reverse();
for task_id in scope_tasks {
self.task_scheduler.cancel(task_id);
}
}
Ok(AsyncExecutionResult::Continue)
}
fn op_emit_event(&mut self) -> Result<AsyncExecutionResult, VMError> {
let (bits, kind) = self.pop_kinded()?;
drop_with_kind(bits, kind);
Ok(AsyncExecutionResult::Continue)
}
}
#[cfg(test)]
pub fn is_async_opcode(opcode: OpCode) -> bool {
matches!(
opcode,
OpCode::Yield
| OpCode::Suspend
| OpCode::Resume
| OpCode::Poll
| OpCode::AwaitBar
| OpCode::AwaitTick
| OpCode::EmitAlert
| OpCode::EmitEvent
| OpCode::Await
| OpCode::SpawnTask
| OpCode::JoinInit
| OpCode::JoinAwait
| OpCode::CancelTask
| OpCode::AsyncScopeEnter
| OpCode::AsyncScopeExit
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_is_async_opcode() {
assert!(is_async_opcode(OpCode::Yield));
assert!(is_async_opcode(OpCode::Suspend));
assert!(is_async_opcode(OpCode::EmitAlert));
assert!(is_async_opcode(OpCode::AsyncScopeEnter));
assert!(is_async_opcode(OpCode::AsyncScopeExit));
assert!(!is_async_opcode(OpCode::AddInt));
assert!(!is_async_opcode(OpCode::Jump));
}
#[test]
fn test_is_async_opcode_all_variants() {
assert!(is_async_opcode(OpCode::Yield));
assert!(is_async_opcode(OpCode::Suspend));
assert!(is_async_opcode(OpCode::Resume));
assert!(is_async_opcode(OpCode::Poll));
assert!(is_async_opcode(OpCode::AwaitBar));
assert!(is_async_opcode(OpCode::AwaitTick));
assert!(is_async_opcode(OpCode::EmitAlert));
assert!(is_async_opcode(OpCode::EmitEvent));
assert!(!is_async_opcode(OpCode::PushConst));
assert!(!is_async_opcode(OpCode::Return));
assert!(!is_async_opcode(OpCode::Call));
assert!(!is_async_opcode(OpCode::Nop));
}
#[test]
fn test_async_execution_result_variants() {
let continue_result = AsyncExecutionResult::Continue;
assert!(matches!(continue_result, AsyncExecutionResult::Continue));
let yielded_result = AsyncExecutionResult::Yielded;
assert!(matches!(yielded_result, AsyncExecutionResult::Yielded));
let suspended_result = AsyncExecutionResult::Suspended(SuspensionInfo {
wait_type: WaitType::AnyEvent,
resume_ip: 42,
});
match suspended_result {
AsyncExecutionResult::Suspended(info) => {
assert_eq!(info.resume_ip, 42);
assert!(matches!(info.wait_type, WaitType::AnyEvent));
}
_ => panic!("Expected Suspended"),
}
}
#[test]
fn test_wait_type_variants() {
let next_bar = WaitType::NextBar {
source: "market_data".to_string(),
};
match next_bar {
WaitType::NextBar { source } => assert_eq!(source, "market_data"),
_ => panic!("Expected NextBar"),
}
let timer = WaitType::Timer { id: 123 };
match timer {
WaitType::Timer { id } => assert_eq!(id, 123),
_ => panic!("Expected Timer"),
}
let any = WaitType::AnyEvent;
assert!(matches!(any, WaitType::AnyEvent));
}
#[test]
fn test_suspension_info_creation() {
let info = SuspensionInfo {
wait_type: WaitType::Timer { id: 999 },
resume_ip: 100,
};
assert_eq!(info.resume_ip, 100);
assert!(matches!(info.wait_type, WaitType::Timer { id: 999 }));
}
#[test]
fn test_is_async_opcode_await() {
assert!(is_async_opcode(OpCode::Await));
}
#[test]
fn test_wait_type_future() {
let future = WaitType::Future { id: 42 };
match future {
WaitType::Future { id } => assert_eq!(id, 42),
_ => panic!("Expected Future"),
}
}
#[test]
fn test_is_async_opcode_join_opcodes() {
assert!(is_async_opcode(OpCode::SpawnTask));
assert!(is_async_opcode(OpCode::JoinInit));
assert!(is_async_opcode(OpCode::JoinAwait));
assert!(is_async_opcode(OpCode::CancelTask));
}
#[test]
fn test_wait_type_task_group() {
let tg = WaitType::TaskGroup {
kind: 0,
task_ids: vec![1, 2, 3],
};
match tg {
WaitType::TaskGroup { kind, task_ids } => {
assert_eq!(kind, 0); assert_eq!(task_ids.len(), 3);
assert_eq!(task_ids, vec![1, 2, 3]);
}
_ => panic!("Expected TaskGroup"),
}
}
#[test]
fn test_wait_type_task_group_race() {
let tg = WaitType::TaskGroup {
kind: 1,
task_ids: vec![10, 20],
};
match tg {
WaitType::TaskGroup { kind, task_ids } => {
assert_eq!(kind, 1); assert_eq!(task_ids, vec![10, 20]);
}
_ => panic!("Expected TaskGroup"),
}
}
}