use crate::runtime::{Coro, CoroStatus, Gc, Table, Value};
use crate::version::LuaVersion;
use crate::vm::argcheck::{Args, check_function, type_error};
use crate::vm::builtins::{arg_error, raise_str};
use crate::vm::error::LuaError;
use crate::vm::exec::Vm;
pub(crate) fn open_coroutine(vm: &mut Vm) {
let t = vm.heap.new_table();
let set = |vm: &mut Vm, name: &str, f: crate::runtime::value::NativeFn| {
let k = Value::Str(vm.heap.intern(name.as_bytes()));
let fv = vm.native(f);
unsafe { t.as_mut() }
.set(&mut vm.heap, k, fv)
.expect("valid key");
};
set(vm, "create", co_create);
set(vm, "yield", co_yield);
set(vm, "status", co_status);
set(vm, "running", co_running);
let in_wrap = Value::Table(vm.heap.new_table());
for (name, f) in [
("resume", co_resume as crate::runtime::value::NativeFn),
("wrap", co_wrap),
] {
let k = Value::Str(vm.heap.intern(name.as_bytes()));
let fv = vm.native_with(f, Box::new([in_wrap]));
unsafe { t.as_mut() }
.set(&mut vm.heap, k, fv)
.expect("valid key");
}
if vm.version() >= LuaVersion::Lua53 {
set(vm, "isyieldable", co_isyieldable);
}
if vm.version() >= LuaVersion::Lua54 {
set(vm, "close", co_close);
}
vm.set_global("coroutine", Value::Table(t))
.expect("stdlib registration");
vm.barrier_back_table(t);
}
fn collect_args(vm: &Vm, fs: u32, nargs: u32) -> Vec<Value> {
(0..nargs).map(|i| vm.nat_arg(fs, nargs, i)).collect()
}
fn check_body(vm: &mut Vm, a: Args) -> Result<Value, LuaError> {
let v = a.get(vm, 0);
if vm.version() <= LuaVersion::Lua51 {
return match v {
Value::Closure(_) if !a.is_none(0) => Ok(v),
_ => Err(arg_error(vm, 1, "Lua function expected")),
};
}
check_function(vm, a, 0)
}
fn co_create(vm: &mut Vm, fs: u32, nargs: u32) -> Result<u32, LuaError> {
let body = check_body(vm, Args::new(fs, nargs))?;
let co = vm.new_coro(body);
Ok(vm.nat_return(fs, &[Value::Coro(co)]))
}
fn check_co(vm: &mut Vm, a: Args) -> Result<Gc<Coro>, LuaError> {
match a.get(vm, 0) {
Value::Coro(co) if !a.is_none(0) => Ok(co),
_ => Err(match vm.version() {
LuaVersion::Lua51 | LuaVersion::Lua52 => arg_error(vm, 1, "coroutine expected"),
LuaVersion::Lua53 => arg_error(vm, 1, "thread expected"),
_ => type_error(vm, a, 0, "thread"),
}),
}
}
fn resume_refusal(
vm: &Vm,
co: Gc<Coro>,
own_frame_empty: bool,
in_wrap: Gc<Table>,
) -> Option<String> {
let status = vm.effective_coro_status(co);
if status == CoroStatus::Suspended {
return None;
}
let empty = if vm.current_coro().is_some_and(|c| c.ptr_eq(co)) {
own_frame_empty
} else {
in_wrap.get(Value::Coro(co)).truthy()
};
Some(match vm.version() {
LuaVersion::Lua51 => format!("cannot resume {} coroutine", vm.coro_status_str(co)),
LuaVersion::Lua52 | LuaVersion::Lua53 if status == CoroStatus::Dead || empty => {
"cannot resume dead coroutine".to_string()
}
_ if status == CoroStatus::Dead => "cannot resume dead coroutine".to_string(),
_ => "cannot resume non-suspended coroutine".to_string(),
})
}
fn upval_table(vm: &Vm, fs: u32, i: usize) -> Gc<Table> {
match vm.nat_upval(fs, i) {
Value::Table(t) => t,
_ => unreachable!("coroutine natives keep their state table here"),
}
}
fn mark_in_wrap(vm: &mut Vm, in_wrap: Gc<Table>, on: bool) -> Result<(), LuaError> {
if !matches!(vm.version(), LuaVersion::Lua52 | LuaVersion::Lua53) {
return Ok(());
}
let me = vm.running_thread().0;
let v = if on { Value::Bool(true) } else { Value::Nil };
if unsafe { in_wrap.as_mut() }
.set(&mut vm.heap, me, v)
.is_err()
{
return Err(vm.rt_err("table overflow"));
}
vm.barrier_back_table(in_wrap);
Ok(())
}
fn co_resume(vm: &mut Vm, fs: u32, nargs: u32) -> Result<u32, LuaError> {
let co = check_co(vm, Args::new(fs, nargs))?;
let in_wrap = upval_table(vm, fs, 0);
if let Some(msg) = resume_refusal(vm, co, false, in_wrap) {
let m = Value::Str(vm.heap.intern(msg.as_bytes()));
return Ok(vm.nat_return(fs, &[Value::Bool(false), m]));
}
let args: Vec<Value> = (1..nargs).map(|i| vm.nat_arg(fs, nargs, i)).collect();
match vm.resume_coro(co, args) {
Ok(mut vals) => {
if (vals.len() as i64) + 1 > vm.stack_room() {
let msg = vm.heap.intern(b"too many results to resume");
return Ok(vm.nat_return(fs, &[Value::Bool(false), Value::Str(msg)]));
}
let mut out = Vec::with_capacity(vals.len() + 1);
out.push(Value::Bool(true));
out.append(&mut vals);
Ok(vm.nat_return(fs, &out))
}
Err(e) => {
let e = death_value(vm, e.0);
Ok(vm.nat_return(fs, &[Value::Bool(false), e]))
}
}
}
fn death_value(vm: &mut Vm, e: Value) -> Value {
if e.is_nil() && vm.version() >= LuaVersion::Lua55 {
return Value::Str(vm.heap.intern(b"<no error object>"));
}
e
}
fn co_yield(vm: &mut Vm, fs: u32, nargs: u32) -> Result<u32, LuaError> {
if let Some(msg) = vm.yield_barrier() {
return Err(vm.plain_err(msg));
}
let vals = collect_args(vm, fs, nargs);
Err(vm.do_yield(fs, vals))
}
fn co_status(vm: &mut Vm, fs: u32, nargs: u32) -> Result<u32, LuaError> {
let co = check_co(vm, Args::new(fs, nargs))?;
let s = vm.coro_status_str(co);
let v = Value::Str(vm.heap.intern(s.as_bytes()));
Ok(vm.nat_return(fs, &[v]))
}
fn co_running(vm: &mut Vm, fs: u32, _nargs: u32) -> Result<u32, LuaError> {
let (thread, is_main) = vm.running_thread();
if vm.version() <= LuaVersion::Lua51 {
let v = if is_main { Value::Nil } else { thread };
return Ok(vm.nat_return(fs, &[v]));
}
Ok(vm.nat_return(fs, &[thread, Value::Bool(is_main)]))
}
fn co_isyieldable(vm: &mut Vm, fs: u32, nargs: u32) -> Result<u32, LuaError> {
let a = Args::new(fs, nargs);
let co = if vm.version() >= LuaVersion::Lua54 && !a.is_none(0) {
Some(check_co(vm, a)?)
} else {
None
};
let y = vm.is_yieldable(co);
Ok(vm.nat_return(fs, &[Value::Bool(y)]))
}
fn co_wrapped(vm: &mut Vm, fs: u32, nargs: u32) -> Result<u32, LuaError> {
let Value::Coro(co) = vm.nat_upval(fs, 0) else {
unreachable!("wrap upvalue is a coroutine");
};
let in_wrap = upval_table(vm, fs, 1);
let err = match resume_refusal(vm, co, nargs == 0, in_wrap) {
Some(msg) => Value::Str(vm.heap.intern(msg.as_bytes())),
None => {
let args = collect_args(vm, fs, nargs);
mark_in_wrap(vm, in_wrap, true)?;
let r = vm.resume_coro(co, args);
mark_in_wrap(vm, in_wrap, false)?;
match r {
Ok(vals) => return Ok(vm.nat_return(fs, &vals)),
Err(_) if vm.version() >= LuaVersion::Lua54 => match vm.close_coro(co) {
Ok(Some(e)) => death_value(vm, e),
Ok(None) => unreachable!("a coroutine that died by error has an error"),
Err(e) => death_value(vm, e.0),
},
Err(e) => death_value(vm, e.0),
}
}
};
Err(LuaError(wrap_where(vm, err)))
}
fn wrap_where(vm: &mut Vm, err: Value) -> Value {
let text = match err {
Value::Str(s) => s.as_bytes().to_vec(),
Value::Int(_) | Value::Float(_) if vm.version() <= LuaVersion::Lua52 => {
crate::vm::argcheck::to_str_bytes(vm, err).expect("a number converts")
}
_ => return err,
};
let mut out = vm.position_prefix().unwrap_or_default().into_bytes();
out.extend_from_slice(&text);
Value::Str(vm.heap.intern(&out))
}
fn co_wrap(vm: &mut Vm, fs: u32, nargs: u32) -> Result<u32, LuaError> {
let body = check_body(vm, Args::new(fs, nargs))?;
let co = vm.new_coro(body);
let in_wrap = vm.nat_upval(fs, 0);
let f = vm.native_with(co_wrapped, Box::new([Value::Coro(co), in_wrap]));
Ok(vm.nat_return(fs, &[f]))
}
fn co_close(vm: &mut Vm, fs: u32, nargs: u32) -> Result<u32, LuaError> {
let a = Args::new(fs, nargs);
let co = if vm.version() >= LuaVersion::Lua55 && a.is_none(0) {
match vm.current_coro() {
Some(c) => c,
None => return Err(raise_str(vm, "cannot close main thread")),
}
} else {
check_co(vm, a)?
};
if vm.version() < LuaVersion::Lua55 && vm.current_coro().is_some_and(|c| c.ptr_eq(co)) {
return Err(raise_str(vm, "cannot close a running coroutine"));
}
match vm.effective_coro_status(co) {
CoroStatus::Dead | CoroStatus::Suspended => match vm.close_coro(co) {
Ok(Some(e)) => {
let e = death_value(vm, e);
Ok(vm.nat_return(fs, &[Value::Bool(false), e]))
}
Ok(None) => Ok(vm.nat_return(fs, &[Value::Bool(true)])),
Err(e) => {
let e = death_value(vm, e.0);
Ok(vm.nat_return(fs, &[Value::Bool(false), e]))
}
},
CoroStatus::Normal => Err(raise_str(vm, "cannot close a normal coroutine")),
CoroStatus::Running => {
if vm.version() >= LuaVersion::Lua55 {
if vm.is_main_coro(co) {
return Err(raise_str(vm, "cannot close main thread"));
}
if vm.current_coro().is_some_and(|c| c.ptr_eq(co)) {
return Err(vm.close_running());
}
}
Err(raise_str(vm, "cannot close a running coroutine"))
}
}
}