#![allow(clippy::needless_doctest_main)]
pub use actor_attribute_macro::actor;
pub mod tls;
pub use tls::TlsConfig;
pub trait Actor {
fn dispatch(
&self,
method_name: &str,
msg: &str,
) -> impl std::future::Future<Output = String> + Send;
fn create_options(self, port: u16, websocket: bool, tls_config: Option<TlsConfig>)
where
Self: Send + Sync + Sized + 'static,
{
let actor = std::sync::Arc::new(self);
let handle = tokio::runtime::Handle::current();
handle.spawn(async move {
match (websocket, tls_config) {
(true, Some(tls_config)) => {
start_websocket_server_with_tls(actor, port, tls_config).await;
}
(true, None) => {
start_websocket_server(actor, port).await;
}
(false, Some(tls_config)) => {
start_http_server_with_tls(actor, port, tls_config).await;
}
(false, None) => {
start_http_server(actor, port).await;
}
}
});
}
fn create(self, port: u16)
where
Self: Send + Sync + Sized + 'static,
{
self.create_options(port, false, None);
}
fn create_ws(self, port: u16)
where
Self: Send + Sync + Sized + 'static,
{
self.create_options(port, true, None);
}
fn create_https(self, port: u16, tls_config: TlsConfig)
where
Self: Send + Sync + Sized + 'static,
{
self.create_options(port, false, Some(tls_config));
}
fn create_wss(self, port: u16, tls_config: TlsConfig)
where
Self: Send + Sync + Sized + 'static,
{
self.create_options(port, true, Some(tls_config));
}
}
use futures_util::{SinkExt, StreamExt};
use http_body_util::Full;
use hyper::body::Bytes;
use hyper::service::service_fn;
use hyper::{Request, Response, StatusCode};
use hyper_util::rt::TokioIo;
use hyper_util::server::conn::auto::Builder;
use std::convert::Infallible;
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::net::TcpListener;
use tokio_tungstenite::{accept_async, tungstenite::Message};
async fn start_http_server<T>(actor: Arc<T>, port: u16)
where
T: Actor + Send + Sync + 'static,
{
let addr = SocketAddr::from(([0, 0, 0, 0], port));
let listener = TcpListener::bind(&addr).await.unwrap_or_else(|e| {
panic!("Failed to bind HTTP server to {addr:?}: {}", e);
});
log::info!("HTTP server listening on http://{}", addr);
loop {
let (stream, _) = match listener.accept().await {
Ok(conn) => conn,
Err(e) => {
log::error!("Failed to accept connection: {}", e);
continue;
}
};
let actor = Arc::clone(&actor);
tokio::spawn(async move {
let io = TokioIo::new(stream);
let service = service_fn(move |req| {
let actor = Arc::clone(&actor);
async move { handle_http_request(actor, req).await }
});
if let Err(e) = Builder::new(hyper_util::rt::TokioExecutor::new())
.serve_connection(io, service)
.await
{
log::error!("HTTP connection error: {}", e);
}
});
}
}
async fn handle_http_request<T>(
actor: Arc<T>,
req: Request<hyper::body::Incoming>,
) -> Result<Response<Full<Bytes>>, Infallible>
where
T: Actor + Send + Sync + 'static,
{
let method = req.method().as_str().to_string();
let path = req.uri().path().to_string();
let body_str = match http_body_util::BodyExt::collect(req.into_body()).await {
Ok(collected) => match std::str::from_utf8(&collected.to_bytes()) {
Ok(s) => s.to_string(),
Err(_) => {
return Ok(Response::builder()
.status(StatusCode::BAD_REQUEST)
.body(Full::new(Bytes::from("Invalid UTF-8 in request body")))
.unwrap());
}
},
Err(_) => {
return Ok(Response::builder()
.status(StatusCode::BAD_REQUEST)
.body(Full::new(Bytes::from("Failed to read request body")))
.unwrap());
}
};
if method == "POST" {
let method_name = path.trim_start_matches('/');
let response_body = (*actor).dispatch(method_name, &body_str).await;
Ok(Response::builder()
.status(StatusCode::OK)
.header("Content-Type", "application/json")
.header("Access-Control-Allow-Origin", "*")
.header("Access-Control-Allow-Methods", "POST, OPTIONS")
.header("Access-Control-Allow-Headers", "Content-Type")
.body(Full::new(Bytes::from(response_body)))
.unwrap())
} else if method == "OPTIONS" {
Ok(Response::builder()
.status(StatusCode::OK)
.header("Access-Control-Allow-Origin", "*")
.header("Access-Control-Allow-Methods", "POST, OPTIONS")
.header("Access-Control-Allow-Headers", "Content-Type")
.header("Content-Length", "0")
.body(Full::new(Bytes::new()))
.unwrap())
} else {
Ok(Response::builder()
.status(StatusCode::METHOD_NOT_ALLOWED)
.header("Content-Type", "text/plain")
.header("Access-Control-Allow-Origin", "*")
.body(Full::new(Bytes::from("Method Not Allowed")))
.unwrap())
}
}
async fn start_websocket_server<T>(actor: Arc<T>, port: u16)
where
T: Actor + Send + Sync + 'static,
{
let addr = SocketAddr::from(([0, 0, 0, 0], port));
let listener = tokio::net::TcpListener::bind(&addr)
.await
.unwrap_or_else(|e| {
panic!(
"Failed to bind WebSocket server with address {addr:?}: {}",
e
);
});
log::info!("WebSocket server listening on ws://{}", addr);
loop {
let (stream, _) = match listener.accept().await {
Ok(conn) => conn,
Err(e) => {
log::error!("Failed to accept WebSocket connection: {}", e);
continue;
}
};
let actor = Arc::clone(&actor);
tokio::spawn(async move {
if let Err(e) = handle_websocket_connection(actor, stream).await {
log::error!("WebSocket connection error: {}", e);
}
});
}
}
async fn handle_websocket_connection<T, S>(
actor: Arc<T>,
stream: S,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>>
where
T: Actor + Send + Sync + 'static,
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let ws_stream = accept_async(stream).await?;
let (mut ws_sender, mut ws_receiver) = ws_stream.split();
while let Some(msg) = ws_receiver.next().await {
match msg? {
Message::Text(text) => {
match serde_json::from_str::<serde_json::Value>(&text) {
Ok(json) => {
if let (Some(method), Some(params)) = (
json.get("method").and_then(|v| v.as_str()),
json.get("params"),
) {
let params_str = params.to_string();
let response = (*actor).dispatch(method, ¶ms_str).await;
if let Err(_e) = ws_sender.send(Message::Text(response)).await {
log::error!("Failed to send WebSocket response: {}", _e);
break;
}
} else {
let error_response = serde_json::json!({
"error": "Invalid message format. Expected {\"method\": \"method_name\", \"params\": {...}}"
}).to_string();
if let Err(e) = ws_sender.send(Message::Text(error_response)).await {
log::error!("Failed to send WebSocket error response: {}", e);
break;
}
}
}
Err(e) => {
let error_response =
serde_json::json!({"error": format!("JSON parse error: {}", e)})
.to_string();
if let Err(e) = ws_sender.send(Message::Text(error_response)).await {
log::error!("Failed to send WebSocket error response: {}", e);
break;
}
}
}
}
Message::Close(_) => {
break;
}
_ => {
}
}
}
Ok(())
}
async fn start_http_server_with_tls<T>(actor: Arc<T>, port: u16, tls_config: TlsConfig)
where
T: Actor + Send + Sync + 'static,
{
match tls_config.load_server_config().await {
Ok(tls_server_config) => {
let addr = SocketAddr::from(([0, 0, 0, 0], port));
let listener = tokio::net::TcpListener::bind(&addr).await.unwrap();
let tls_acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(tls_server_config));
log::info!("HTTPS server listening on https://{}", addr);
loop {
let (stream, _) = match listener.accept().await {
Ok(conn) => conn,
Err(e) => {
log::error!("Failed to accept HTTPS connection: {}", e);
continue;
}
};
let actor = Arc::clone(&actor);
let tls_acceptor = tls_acceptor.clone();
tokio::spawn(async move {
match tls_acceptor.accept(stream).await {
Ok(tls_stream) => {
if let Err(e) = handle_https_connection(actor, tls_stream).await {
log::error!("HTTPS connection error: {}", e);
}
}
Err(e) => {
log::error!("TLS handshake error: {}", e);
}
}
});
}
}
Err(e) => {
log::error!("Failed to load TLS configuration: {}", e);
}
}
}
async fn start_websocket_server_with_tls<T>(actor: Arc<T>, port: u16, tls_config: TlsConfig)
where
T: Actor + Send + Sync + 'static,
{
let addr = SocketAddr::from(([0, 0, 0, 0], port));
let listener = tokio::net::TcpListener::bind(&addr)
.await
.unwrap_or_else(|e| {
panic!("Failed to bind WSS server address {addr:?}: {}", e);
});
match tls_config.load_server_config().await {
Ok(tls_server_config) => {
let tls_acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(tls_server_config));
log::info!("WSS server listening on wss://{}", addr);
loop {
let (stream, _) = match listener.accept().await {
Ok(conn) => conn,
Err(e) => {
log::error!("Failed to accept WSS connection: {}", e);
continue;
}
};
let actor = Arc::clone(&actor);
let tls_acceptor = tls_acceptor.clone();
tokio::spawn(async move {
match tls_acceptor.accept(stream).await {
Ok(tls_stream) => {
if let Err(e) = handle_websocket_connection(actor, tls_stream).await {
log::error!("WSS connection error: {}", e);
}
}
Err(e) => {
log::error!("TLS handshake error: {}", e);
}
}
});
}
}
Err(e) => {
log::error!("Failed to load TLS configuration: {}", e);
}
}
}
async fn handle_https_connection<T>(
actor: Arc<T>,
stream: tokio_rustls::server::TlsStream<tokio::net::TcpStream>,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>>
where
T: Actor + Send + Sync + 'static,
{
let io = TokioIo::new(stream);
let service = service_fn(move |req| {
let actor = actor.clone();
async move { handle_http_request(actor, req).await }
});
if let Err(e) = Builder::new(hyper_util::rt::TokioExecutor::new())
.serve_connection(io, service)
.await
{
log::error!("HTTPS connection error: {}", e);
}
Ok(())
}
#[cfg(test)]
mod test_actor;