use crate::async_task::AsyncTask;
use crate::error::IsleError;
use crate::hook::{self, CancelToken};
use crate::thread;
use std::thread::JoinHandle;
type AsyncExecFn = Box<dyn FnOnce(&mlua::Lua) -> Result<String, IsleError> + Send>;
type AsyncResultTx = tokio::sync::oneshot::Sender<Result<String, IsleError>>;
enum AsyncRequest {
Eval {
code: String,
cancel: CancelToken,
tx: AsyncResultTx,
},
Call {
func: String,
args: Vec<String>,
cancel: CancelToken,
tx: AsyncResultTx,
},
Exec {
f: AsyncExecFn,
cancel: CancelToken,
tx: AsyncResultTx,
},
CoroutineEval {
code: String,
cancel: CancelToken,
tx: AsyncResultTx,
},
CoroutineCall {
func: String,
args: Vec<String>,
cancel: CancelToken,
tx: AsyncResultTx,
},
Shutdown,
}
const DEFAULT_CHANNEL_CAPACITY: usize = 256;
#[derive(Clone)]
pub struct AsyncIsle {
tx: tokio::sync::mpsc::Sender<AsyncRequest>,
}
#[must_use = "call .shutdown().await for clean thread join; dropping without shutdown detaches the thread"]
pub struct AsyncIsleDriver {
tx: tokio::sync::mpsc::Sender<AsyncRequest>,
done_rx: Option<tokio::sync::oneshot::Receiver<()>>,
join: Option<JoinHandle<()>>,
}
pub struct AsyncIsleBuilder {
channel_capacity: usize,
thread_name: String,
}
impl Default for AsyncIsleBuilder {
fn default() -> Self {
Self {
channel_capacity: DEFAULT_CHANNEL_CAPACITY,
thread_name: "mlua-isle-async".into(),
}
}
}
impl AsyncIsleBuilder {
pub fn channel_capacity(mut self, capacity: usize) -> Self {
self.channel_capacity = capacity;
self
}
pub fn thread_name(mut self, name: &str) -> Self {
self.thread_name = name.to_string();
self
}
pub async fn spawn<F>(self, init: F) -> Result<(AsyncIsle, AsyncIsleDriver), IsleError>
where
F: FnOnce(&mlua::Lua) -> Result<(), mlua::Error> + Send + 'static,
{
AsyncIsle::spawn_inner(init, self.channel_capacity, self.thread_name).await
}
}
impl AsyncIsle {
pub fn builder() -> AsyncIsleBuilder {
AsyncIsleBuilder::default()
}
pub async fn spawn<F>(init: F) -> Result<(Self, AsyncIsleDriver), IsleError>
where
F: FnOnce(&mlua::Lua) -> Result<(), mlua::Error> + Send + 'static,
{
Self::spawn_inner(init, DEFAULT_CHANNEL_CAPACITY, "mlua-isle-async".into()).await
}
async fn spawn_inner<F>(
init: F,
channel_capacity: usize,
thread_name: String,
) -> Result<(Self, AsyncIsleDriver), IsleError>
where
F: FnOnce(&mlua::Lua) -> Result<(), mlua::Error> + Send + 'static,
{
let (tx, rx) = tokio::sync::mpsc::channel::<AsyncRequest>(channel_capacity);
let (init_tx, init_rx) = tokio::sync::oneshot::channel::<Result<(), IsleError>>();
let (done_tx, done_rx) = tokio::sync::oneshot::channel::<()>();
let join = std::thread::Builder::new()
.name(thread_name)
.spawn(move || {
let rt = match tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
{
Ok(rt) => rt,
Err(e) => {
let _ = init_tx.send(Err(IsleError::Init(format!(
"tokio runtime build failed: {e}"
))));
let _ = done_tx.send(());
return;
}
};
let lua = mlua::Lua::new();
match init(&lua) {
Ok(()) => {
let _ = init_tx.send(Ok(()));
run_async_loop(lua, rx, rt);
}
Err(e) => {
let _ = init_tx.send(Err(IsleError::Init(e.to_string())));
}
}
let _ = done_tx.send(());
})
.map_err(|e| IsleError::Init(format!("thread spawn failed: {e}")))?;
init_rx
.await
.map_err(|e| IsleError::Init(format!("init channel closed: {e}")))??;
let handle = Self { tx: tx.clone() };
let driver = AsyncIsleDriver {
tx,
done_rx: Some(done_rx),
join: Some(join),
};
Ok((handle, driver))
}
pub async fn eval(&self, code: &str) -> Result<String, IsleError> {
self.spawn_eval(code).await
}
pub fn spawn_eval(&self, code: &str) -> AsyncTask {
let cancel = CancelToken::new();
let (resp_tx, resp_rx) = tokio::sync::oneshot::channel();
let req = AsyncRequest::Eval {
code: code.to_string(),
cancel: cancel.clone(),
tx: resp_tx,
};
match self.tx.try_send(req) {
Ok(()) => AsyncTask::new(resp_rx, cancel),
Err(e) => make_error_task(try_send_to_isle_error(e), cancel),
}
}
pub async fn call(&self, func: &str, args: &[&str]) -> Result<String, IsleError> {
self.spawn_call(func, args).await
}
pub fn spawn_call(&self, func: &str, args: &[&str]) -> AsyncTask {
let cancel = CancelToken::new();
let (resp_tx, resp_rx) = tokio::sync::oneshot::channel();
let req = AsyncRequest::Call {
func: func.to_string(),
args: args.iter().map(|s| s.to_string()).collect(),
cancel: cancel.clone(),
tx: resp_tx,
};
match self.tx.try_send(req) {
Ok(()) => AsyncTask::new(resp_rx, cancel),
Err(e) => make_error_task(try_send_to_isle_error(e), cancel),
}
}
pub async fn exec<F>(&self, f: F) -> Result<String, IsleError>
where
F: FnOnce(&mlua::Lua) -> Result<String, IsleError> + Send + 'static,
{
self.spawn_exec(f).await
}
pub fn spawn_exec<F>(&self, f: F) -> AsyncTask
where
F: FnOnce(&mlua::Lua) -> Result<String, IsleError> + Send + 'static,
{
let cancel = CancelToken::new();
let (resp_tx, resp_rx) = tokio::sync::oneshot::channel();
let req = AsyncRequest::Exec {
f: Box::new(f),
cancel: cancel.clone(),
tx: resp_tx,
};
match self.tx.try_send(req) {
Ok(()) => AsyncTask::new(resp_rx, cancel),
Err(e) => make_error_task(try_send_to_isle_error(e), cancel),
}
}
pub async fn coroutine_eval(&self, code: &str) -> Result<String, IsleError> {
self.spawn_coroutine_eval(code).await
}
pub fn spawn_coroutine_eval(&self, code: &str) -> AsyncTask {
let cancel = CancelToken::new();
let (resp_tx, resp_rx) = tokio::sync::oneshot::channel();
let req = AsyncRequest::CoroutineEval {
code: code.to_string(),
cancel: cancel.clone(),
tx: resp_tx,
};
match self.tx.try_send(req) {
Ok(()) => AsyncTask::new(resp_rx, cancel),
Err(e) => make_error_task(try_send_to_isle_error(e), cancel),
}
}
pub async fn coroutine_call(&self, func: &str, args: &[&str]) -> Result<String, IsleError> {
self.spawn_coroutine_call(func, args).await
}
pub fn spawn_coroutine_call(&self, func: &str, args: &[&str]) -> AsyncTask {
let cancel = CancelToken::new();
let (resp_tx, resp_rx) = tokio::sync::oneshot::channel();
let req = AsyncRequest::CoroutineCall {
func: func.to_string(),
args: args.iter().map(|s| s.to_string()).collect(),
cancel: cancel.clone(),
tx: resp_tx,
};
match self.tx.try_send(req) {
Ok(()) => AsyncTask::new(resp_rx, cancel),
Err(e) => make_error_task(try_send_to_isle_error(e), cancel),
}
}
pub fn is_alive(&self) -> bool {
!self.tx.is_closed()
}
}
impl AsyncIsleDriver {
pub async fn shutdown(mut self) -> Result<(), IsleError> {
let _ = self.tx.send(AsyncRequest::Shutdown).await;
if let Some(done_rx) = self.done_rx.take() {
done_rx.await.map_err(|_| IsleError::ThreadPanic)?;
}
if let Some(join) = self.join.take() {
join.join().map_err(|_| IsleError::ThreadPanic)?;
}
Ok(())
}
pub fn is_alive(&self) -> bool {
self.join.as_ref().is_some_and(|j| !j.is_finished())
}
}
fn make_error_task(err: IsleError, cancel: CancelToken) -> AsyncTask {
let (tx, rx) = tokio::sync::oneshot::channel();
let _ = tx.send(Err(err));
AsyncTask::new(rx, cancel)
}
fn try_send_to_isle_error<T>(err: tokio::sync::mpsc::error::TrySendError<T>) -> IsleError {
match err {
tokio::sync::mpsc::error::TrySendError::Full(_) => IsleError::ChannelFull,
tokio::sync::mpsc::error::TrySendError::Closed(_) => IsleError::Shutdown,
}
}
fn run_async_loop(
lua: mlua::Lua,
mut rx: tokio::sync::mpsc::Receiver<AsyncRequest>,
rt: tokio::runtime::Runtime,
) {
let local = tokio::task::LocalSet::new();
local.spawn_local(async move {
while let Some(req) = rx.recv().await {
match req {
AsyncRequest::Eval { code, cancel, tx } => {
let result = thread::execute_eval(&lua, &code, &cancel);
let _ = tx.send(result);
}
AsyncRequest::Call {
func,
args,
cancel,
tx,
} => {
let result = thread::execute_call(&lua, &func, &args, &cancel);
let _ = tx.send(result);
}
AsyncRequest::Exec { f, cancel, tx } => {
let result = thread::execute_exec(&lua, f, &cancel);
let _ = tx.send(result);
}
AsyncRequest::CoroutineEval { code, cancel, tx } => {
let lua = lua.clone();
tokio::task::spawn_local(async move {
let result = execute_coroutine_eval(&lua, &code, &cancel).await;
let _ = tx.send(result);
});
}
AsyncRequest::CoroutineCall {
func,
args,
cancel,
tx,
} => {
let lua = lua.clone();
tokio::task::spawn_local(async move {
let result = execute_coroutine_call(&lua, &func, &args, &cancel).await;
let _ = tx.send(result);
});
}
AsyncRequest::Shutdown => break,
}
tokio::task::yield_now().await;
}
});
rt.block_on(local);
}
async fn execute_coroutine_eval(
lua: &mlua::Lua,
code: &str,
cancel: &CancelToken,
) -> Result<String, IsleError> {
let func = lua.load(code).into_function().map_err(IsleError::from)?;
let thread = lua.create_thread(func).map_err(IsleError::from)?;
hook::install_cancel_hook_on_thread(&thread, cancel.clone(), thread::HOOK_INTERVAL)?;
let val: mlua::Value = thread
.into_async(())
.map_err(IsleError::from)?
.await
.map_err(IsleError::from)?;
thread::lua_value_to_string(lua, val)
}
async fn execute_coroutine_call(
lua: &mlua::Lua,
func_name: &str,
args: &[String],
cancel: &CancelToken,
) -> Result<String, IsleError> {
let func: mlua::Function = lua
.globals()
.get(func_name)
.map_err(|e| IsleError::Lua(format!("function '{func_name}' not found: {e}")))?;
let thread = lua.create_thread(func).map_err(IsleError::from)?;
hook::install_cancel_hook_on_thread(&thread, cancel.clone(), thread::HOOK_INTERVAL)?;
let lua_args: Vec<mlua::Value> = args
.iter()
.map(|s| lua.create_string(s).map(mlua::Value::String))
.collect::<mlua::Result<Vec<_>>>()
.map_err(IsleError::from)?;
let multi = mlua::MultiValue::from_vec(lua_args);
let val: mlua::Value = thread
.into_async(multi)
.map_err(IsleError::from)?
.await
.map_err(IsleError::from)?;
thread::lua_value_to_string(lua, val)
}