use std::fmt::Display;
use async_trait::async_trait;
use tokio::sync::{mpsc, oneshot};
pub trait Worker: Send + 'static {
type Request: Send + 'static;
type Response: Send + 'static;
fn handle(&mut self, req: Self::Request) -> Self::Response;
}
#[async_trait]
pub trait AsyncWorker: Send + 'static {
type Request: Send + 'static;
type Response: Send + 'static;
async fn handle(&mut self, req: Self::Request) -> Self::Response;
}
pub struct Socket<Req, Res> {
sender: mpsc::Sender<Task<Req, Res>>,
}
impl<Req, Res> Clone for Socket<Req, Res> {
fn clone(&self) -> Self {
Self {
sender: self.sender.clone(),
}
}
}
pub struct WeakSocket<Req, Res> {
sender: mpsc::WeakSender<Task<Req, Res>>,
}
impl<Req, Res> Clone for WeakSocket<Req, Res> {
fn clone(&self) -> Self {
Self {
sender: self.sender.clone(),
}
}
}
#[derive(Debug, Clone, Copy)]
pub enum RunError<Req> {
FailedToEnqueueReq(Req),
FailedToGetResponse,
}
impl<Req> Display for RunError<Req> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match &self {
RunError::FailedToEnqueueReq(_) => {
write!(f, "failed to enqueue a request.")
},
RunError::FailedToGetResponse => {
write!(f, "failed to recv the response")
},
}
}
}
pub trait Executor {
fn spawn<H: Worker>(handler: H) -> Socket<H::Request, H::Response>;
fn spawn_async<H: AsyncWorker>(handler: H) -> Socket<H::Request, H::Response>;
}
pub struct Task<Req, Res> {
pub request: Req,
respond: Option<oneshot::Sender<Res>>,
}
impl<Req, Res> Task<Req, Res> {
pub fn respond(self, response: Res) {
if let Some(tx) = self.respond {
let _ = tx.send(response);
}
}
}
#[derive(Clone, Copy, Debug)]
pub struct DedicatedThread;
#[derive(Clone, Copy, Debug)]
pub struct TokioSpawn;
impl Executor for DedicatedThread {
fn spawn<H: Worker>(handler: H) -> Socket<H::Request, H::Response> {
let (tx, rx) = mpsc::channel(64);
std::thread::spawn(|| run_blocking(rx, handler));
Socket { sender: tx }
}
fn spawn_async<H: AsyncWorker>(handler: H) -> Socket<H::Request, H::Response> {
let (tx, rx) = mpsc::channel(64);
std::thread::spawn(|| run_blocking_async(rx, handler));
Socket { sender: tx }
}
}
impl Executor for TokioSpawn {
fn spawn<H: Worker>(handler: H) -> Socket<H::Request, H::Response> {
let (tx, rx) = mpsc::channel(64);
tokio::spawn(run_non_blocking(rx, handler));
Socket { sender: tx }
}
fn spawn_async<H: AsyncWorker>(handler: H) -> Socket<H::Request, H::Response> {
let (tx, rx) = mpsc::channel(64);
tokio::spawn(run_non_blocking_async(rx, handler));
Socket { sender: tx }
}
}
fn run_blocking<Req, Res, H: Worker<Request = Req, Response = Res>>(
mut rx: mpsc::Receiver<Task<Req, Res>>,
mut handler: H,
) {
while let Some(event) = rx.blocking_recv() {
let res = handler.handle(event.request);
if let Some(tx) = event.respond {
let _ = tx.send(res);
}
}
}
fn run_blocking_async<Req, Res, H: AsyncWorker<Request = Req, Response = Res>>(
mut rx: mpsc::Receiver<Task<Req, Res>>,
mut handler: H,
) {
while let Some(event) = rx.blocking_recv() {
let res = futures::executor::block_on(handler.handle(event.request));
if let Some(tx) = event.respond {
let _ = tx.send(res);
}
}
}
async fn run_non_blocking<Req, Res, H: Worker<Request = Req, Response = Res>>(
mut rx: mpsc::Receiver<Task<Req, Res>>,
mut handler: H,
) {
while let Some(event) = rx.recv().await {
let result = handler.handle(event.request);
if let Some(tx) = event.respond {
let _ = tx.send(result);
}
}
}
async fn run_non_blocking_async<Req, Res, H: AsyncWorker<Request = Req, Response = Res>>(
mut rx: mpsc::Receiver<Task<Req, Res>>,
mut handler: H,
) {
while let Some(event) = rx.recv().await {
let result = handler.handle(event.request).await;
if let Some(tx) = event.respond {
let _ = tx.send(result);
}
}
}
impl<Req, Res> Socket<Req, Res> {
pub fn raw_bounded(bound: usize) -> (Self, mpsc::Receiver<Task<Req, Res>>) {
let (tx, rx) = mpsc::channel(bound);
let socket = Socket { sender: tx };
(socket, rx)
}
pub fn downgrade(&self) -> WeakSocket<Req, Res> {
WeakSocket {
sender: self.sender.downgrade(),
}
}
pub async fn enqueue(&self, request: Req) -> Result<(), mpsc::error::SendError<Req>> {
let event = Task {
request,
respond: None,
};
self.sender
.send(event)
.await
.map_err(|e| mpsc::error::SendError(e.0.request))
}
pub async fn run(&self, request: Req) -> Result<Res, RunError<Req>> {
let (tx, rx) = oneshot::channel::<Res>();
let event = Task {
request,
respond: Some(tx),
};
self.sender
.send(event)
.await
.map_err(|e| RunError::FailedToEnqueueReq(e.0.request))?;
let result = rx.await.map_err(|_| RunError::FailedToGetResponse)?;
Ok(result)
}
}
impl<Req, Res> WeakSocket<Req, Res> {
pub fn upgrade(&self) -> Option<Socket<Req, Res>> {
self.sender.upgrade().map(|sender| Socket { sender })
}
}
impl<Req, Res> Unpin for Socket<Req, Res> {}
impl<Req, Res> Unpin for WeakSocket<Req, Res> {}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Default)]
struct CounterWorker {
current: u64,
}
impl Worker for CounterWorker {
type Request = u64;
type Response = u64;
fn handle(&mut self, req: Self::Request) -> Self::Response {
self.current += req;
self.current
}
}
#[async_trait]
impl AsyncWorker for CounterWorker {
type Request = u64;
type Response = u64;
async fn handle(&mut self, req: Self::Request) -> Self::Response {
self.current += req;
self.current
}
}
#[tokio::test]
async fn test_dedicated_thread() {
let socket = DedicatedThread::spawn(CounterWorker::default());
assert_eq!(socket.run(10).await.unwrap(), 10);
assert_eq!(socket.run(3).await.unwrap(), 13);
let socket = DedicatedThread::spawn_async(CounterWorker::default());
assert_eq!(socket.run(10).await.unwrap(), 10);
assert_eq!(socket.run(3).await.unwrap(), 13);
}
#[tokio::test]
async fn test_tokio_spawn() {
let socket = TokioSpawn::spawn(CounterWorker::default());
assert_eq!(socket.run(10).await.unwrap(), 10);
assert_eq!(socket.run(3).await.unwrap(), 13);
let socket = DedicatedThread::spawn_async(CounterWorker::default());
assert_eq!(socket.run(10).await.unwrap(), 10);
assert_eq!(socket.run(3).await.unwrap(), 13);
}
}