Refactor websocket

This commit is contained in:
Tobias Reisinger 2026-09-02 00:25:56 +02:00
commit f3d367e479
Signed by: serguzim
GPG key ID: 13AD60C237A28DFE
22 changed files with 924 additions and 773 deletions

View file

@ -1,14 +1,12 @@
use actix::Addr;
use actix_web::{delete, get, put, web, HttpResponse};
use emgauwa_common::db::DbController;
use emgauwa_common::errors::{DatabaseError, EmgauwaError};
use emgauwa_common::models::{convert_db_list, Controller, FromDbModel};
use emgauwa_common::types::{ControllerWsAction, EmgauwaUid, RequestControllerUpdate};
use emgauwa_common::types::{ControllersWsAction, EmgauwaUid, RequestControllerUpdate};
use sqlx::{Pool, Sqlite};
use crate::app_state;
use crate::app_state::AppState;
use tokio::sync::mpsc;
use crate::handlers::EmgauwaMessage;
use crate::server::WsServerAction;
#[get("/controllers")]
pub async fn index(pool: web::Data<Pool<Sqlite>>) -> Result<HttpResponse, EmgauwaError> {
@ -17,7 +15,7 @@ pub async fn index(pool: web::Data<Pool<Sqlite>>) -> Result<HttpResponse, Emgauw
let db_controllers = DbController::get_all(&mut pool_conn).await?;
let controllers: Vec<Controller> = convert_db_list(&mut pool_conn, db_controllers)?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().json(controllers))
}
@ -36,13 +34,14 @@ pub async fn show(
.ok_or(DatabaseError::NotFound)?;
let return_controller = Controller::from_db_model(&mut pool_conn, controller)?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().json(return_controller))
}
#[put("/controllers/{controller_id}")]
pub async fn update(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
path: web::Path<(String,)>,
data: web::Json<RequestControllerUpdate>,
) -> Result<HttpResponse, EmgauwaError> {
@ -60,13 +59,9 @@ pub async fn update(
.await?;
let return_controller = Controller::from_db_model(&mut pool_conn, controller)?;
pool_conn.close().await?;
app_state
.send(app_state::Action {
controller_uid: uid.clone(),
action: ControllerWsAction::Controller(return_controller.clone()),
})
.await??;
ws_server_tx.send(WsServerAction::ForwardToController(uid, ControllersWsAction::Controller(return_controller.clone())))?;
Ok(HttpResponse::Ok().json(return_controller))
}
@ -74,7 +69,7 @@ pub async fn update(
#[delete("/controllers/{controller_id}")]
pub async fn delete(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
path: web::Path<(String,)>,
) -> Result<HttpResponse, EmgauwaError> {
let mut pool_conn = pool.acquire().await?;
@ -82,12 +77,11 @@ pub async fn delete(
let (controller_uid,) = path.into_inner();
let uid = EmgauwaUid::try_from(controller_uid.as_str())?;
app_state
.send(app_state::DisconnectController {
controller_uid: uid.clone(),
})
.await??;
log::debug!("Deleting controller with uid: {}", uid.clone());
DbController::delete_by_uid(&mut pool_conn, uid.clone()).await?;
pool_conn.close().await?;
ws_server_tx.send(WsServerAction::DisconnectController(uid.clone(), false))?;
DbController::delete_by_uid(&mut pool_conn, uid).await?;
Ok(HttpResponse::Ok().emgauwa_message("controller got deleted"))
}

View file

@ -1,17 +1,15 @@
use actix::Addr;
use actix_web::{delete, get, post, put, web, HttpResponse};
use emgauwa_common::db::{DbController, DbMacro};
use emgauwa_common::errors::{DatabaseError, EmgauwaError};
use emgauwa_common::models::{convert_db_list, FromDbModel, Macro, MacroAction, Relay};
use emgauwa_common::types::{
ControllerWsAction, EmgauwaUid, RequestMacroCreate, RequestMacroExecute, RequestMacroUpdate,
ControllersWsAction, EmgauwaUid, RequestMacroCreate, RequestMacroExecute, RequestMacroUpdate,
};
use sqlx::pool::PoolConnection;
use sqlx::{Pool, Sqlite};
use crate::app_state;
use crate::app_state::AppState;
use tokio::sync::mpsc;
use crate::handlers::EmgauwaMessage;
use crate::server::WsServerAction;
#[get("/macros")]
pub async fn index(pool: web::Data<Pool<Sqlite>>) -> Result<HttpResponse, EmgauwaError> {
@ -19,7 +17,7 @@ pub async fn index(pool: web::Data<Pool<Sqlite>>) -> Result<HttpResponse, Emgauw
let db_macros = DbMacro::get_all(&mut pool_conn).await?;
let macros: Vec<Macro> = convert_db_list(&mut pool_conn, db_macros)?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().json(macros))
}
@ -38,6 +36,7 @@ pub async fn show(
.ok_or(DatabaseError::NotFound)?;
let return_macro = Macro::from_db_model(&mut pool_conn, db_macro)?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().json(return_macro))
}
@ -55,6 +54,7 @@ pub async fn add(
.await?;
let return_macro = Macro::from_db_model(&mut pool_conn, new_macro)?;
pool_conn.close().await?;
Ok(HttpResponse::Created().json(return_macro))
}
@ -84,6 +84,7 @@ pub async fn update(
}
let return_macro = Macro::from_db_model(&mut pool_conn, db_macro)?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().json(return_macro))
}
@ -98,14 +99,15 @@ pub async fn delete(
let uid = EmgauwaUid::try_from(macro_uid.as_str())?;
DbMacro::delete_by_uid(&mut pool_conn, uid).await?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().emgauwa_message("macro got deleted"))
}
#[put("/macros/{macro_id}/execute")]
pub async fn execute(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
path: web::Path<(String,)>,
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
query: web::Query<RequestMacroExecute>,
) -> Result<HttpResponse, EmgauwaError> {
let mut pool_conn = pool.acquire().await?;
@ -137,14 +139,10 @@ pub async fn execute(
let affected_relays =
collect_affected_relays(&mut pool_conn, &mut actions, &controller).await?;
app_state
.send(app_state::Action {
controller_uid: controller.uid,
action: ControllerWsAction::Relays(affected_relays.clone()),
})
.await??;
ws_server_tx.send(WsServerAction::ForwardToController(uid.clone(), ControllersWsAction::Relays(affected_relays.clone())))?;
}
pool_conn.close().await?;
Ok(HttpResponse::Ok().emgauwa_message("macro got executed"))
}

View file

@ -1,42 +1,55 @@
use actix::Addr;
use crate::handlers::EmgauwaMessage;
use crate::server::WsServerAction;
use actix_web::{get, post, put, web, HttpResponse};
use emgauwa_common::db::{DbController, DbJunctionRelaySchedule, DbRelay, DbTag};
use emgauwa_common::errors::{DatabaseError, EmgauwaError};
use emgauwa_common::models::{convert_db_list, FromDbModel, Relay};
use emgauwa_common::types::{ControllerWsAction, EmgauwaUid, RequestRelayPulse, RequestRelayUpdate};
use emgauwa_common::types::{
ControllersWsAction, EmgauwaUid, RequestRelayPulse, RequestRelayUpdate,
};
use sqlx::{Pool, Sqlite};
use tokio::sync::{mpsc, oneshot};
use crate::app_state;
use crate::app_state::AppState;
use crate::handlers::EmgauwaMessage;
pub async fn get_stated_relays(app_state: &Addr<AppState>) -> Result<Vec<Relay>, EmgauwaError> {
app_state.send(app_state::GetRelays {}).await?
pub async fn get_stated_relays(
ws_server_tx: &web::Data<mpsc::UnboundedSender<WsServerAction>>,
) -> Result<Vec<Relay>, EmgauwaError> {
let (res_tx, res_rx) = oneshot::channel();
ws_server_tx.send(WsServerAction::GetRelays(res_tx))?;
res_rx.await?
}
pub async fn load_state_for_relay(relay: &mut Relay, app_state: &Addr<AppState>) -> Result<(), EmgauwaError>{
let stated_relays = get_stated_relays(app_state).await?;
pub async fn load_state_for_relay(
relay: &mut Relay,
ws_server_tx: &web::Data<mpsc::UnboundedSender<WsServerAction>>,
) -> Result<(), EmgauwaError> {
let stated_relays = get_stated_relays(ws_server_tx).await?;
relay.find_and_apply_state(&stated_relays);
Ok(())
}
pub async fn load_state_for_relays(relays: &mut [Relay], app_state: &Addr<AppState>) -> Result<(), EmgauwaError>{
let stated_relays = get_stated_relays(app_state).await?;
relays.iter_mut().for_each(|r| r.find_and_apply_state(&stated_relays));
pub async fn load_state_for_relays(
relays: &mut [Relay],
ws_server_tx: &web::Data<mpsc::UnboundedSender<WsServerAction>>,
) -> Result<(), EmgauwaError> {
let stated_relays = get_stated_relays(ws_server_tx).await?;
relays
.iter_mut()
.for_each(|r| r.find_and_apply_state(&stated_relays));
Ok(())
}
#[get("/relays")]
pub async fn index(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
) -> Result<HttpResponse, EmgauwaError> {
let mut pool_conn = pool.acquire().await?;
let db_relays = DbRelay::get_all(&mut pool_conn).await?;
let mut relays: Vec<Relay> = convert_db_list(&mut pool_conn, db_relays)?;
load_state_for_relays(&mut relays, &app_state).await?;
pool_conn.close().await?;
load_state_for_relays(&mut relays, &ws_server_tx).await?;
Ok(HttpResponse::Ok().json(relays))
}
@ -44,7 +57,7 @@ pub async fn index(
#[get("/relays/tag/{tag}")]
pub async fn tagged(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
path: web::Path<(String,)>,
) -> Result<HttpResponse, EmgauwaError> {
let mut pool_conn = pool.acquire().await?;
@ -57,7 +70,8 @@ pub async fn tagged(
let db_relays = DbRelay::get_by_tag(&mut pool_conn, &tag_db).await?;
let mut relays: Vec<Relay> = convert_db_list(&mut pool_conn, db_relays)?;
load_state_for_relays(&mut relays, &app_state).await?;
pool_conn.close().await?;
load_state_for_relays(&mut relays, &ws_server_tx).await?;
Ok(HttpResponse::Ok().json(relays))
}
@ -65,7 +79,7 @@ pub async fn tagged(
#[get("/controllers/{controller_id}/relays")]
pub async fn index_for_controller(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
path: web::Path<(String,)>,
) -> Result<HttpResponse, EmgauwaError> {
let mut pool_conn = pool.acquire().await?;
@ -80,7 +94,8 @@ pub async fn index_for_controller(
let db_relays = controller.get_relays(&mut pool_conn).await?;
let mut relays: Vec<Relay> = convert_db_list(&mut pool_conn, db_relays)?;
load_state_for_relays(&mut relays, &app_state).await?;
pool_conn.close().await?;
load_state_for_relays(&mut relays, &ws_server_tx).await?;
Ok(HttpResponse::Ok().json(relays))
}
@ -88,7 +103,7 @@ pub async fn index_for_controller(
#[get("/controllers/{controller_id}/relays/{relay_num}")]
pub async fn show_for_controller(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
path: web::Path<(String, i64)>,
) -> Result<HttpResponse, EmgauwaError> {
let mut pool_conn = pool.acquire().await?;
@ -105,7 +120,8 @@ pub async fn show_for_controller(
.ok_or(DatabaseError::NotFound)?;
let mut relay = Relay::from_db_model(&mut pool_conn, db_relay)?;
load_state_for_relay(&mut relay, &app_state).await?;
pool_conn.close().await?;
load_state_for_relay(&mut relay, &ws_server_tx).await?;
Ok(HttpResponse::Ok().json(relay))
}
@ -113,7 +129,7 @@ pub async fn show_for_controller(
#[put("/controllers/{controller_id}/relays/{relay_num}")]
pub async fn update_for_controller(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
path: web::Path<(String, i64)>,
data: web::Json<RequestRelayUpdate>,
) -> Result<HttpResponse, EmgauwaError> {
@ -158,24 +174,25 @@ pub async fn update_for_controller(
let mut return_relay = Relay::from_db_model(&mut pool_conn, relay)?;
load_state_for_relay(&mut return_relay, &app_state).await?;
load_state_for_relay(&mut return_relay, &ws_server_tx).await?;
match &data.override_schedule {
Some(Some(s_uid)) => { // We want to set an override schedule
Some(Some(s_uid)) => {
// We want to set an override schedule
let schedule = s_uid.get_schedule(&mut pool_conn).await?;
return_relay.override_schedule = Some(schedule);
}
Some(None) => { // We want to unset the override schedule
Some(None) => {
// We want to unset the override schedule
return_relay.override_schedule = None;
}
None => {} // We want to keep the override schedule as is
}
app_state
.send(app_state::Action {
controller_uid: uid,
action: ControllerWsAction::Relays(vec![return_relay.clone()]),
})
.await??;
pool_conn.close().await?;
ws_server_tx.send(WsServerAction::ForwardToController(
uid,
ControllersWsAction::Relays(vec![return_relay.clone()]),
))?;
Ok(HttpResponse::Ok().json(return_relay))
}
@ -183,7 +200,7 @@ pub async fn update_for_controller(
#[post("/controllers/{controller_id}/relays/{relay_num}/pulse")]
pub async fn pulse(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
path: web::Path<(String, i64)>,
data: web::Json<RequestRelayPulse>,
) -> Result<HttpResponse, EmgauwaError> {
@ -200,14 +217,13 @@ pub async fn pulse(
.await?
.ok_or(DatabaseError::NotFound)?;
let duration = data.duration.filter(|&d| d > 0);
pool_conn.close().await?;
app_state
.send(app_state::Action {
controller_uid: uid,
action: ControllerWsAction::RelayPulse((relay.number, duration)),
})
.await??;
let duration = data.duration.filter(|&d| d > 0);
ws_server_tx.send(WsServerAction::ForwardToController(
uid,
ControllersWsAction::RelayPulse((relay.number, duration)),
))?;
Ok(HttpResponse::Ok().emgauwa_message("pulse sent"))
}

View file

@ -1,19 +1,17 @@
use actix::Addr;
use crate::handlers::EmgauwaMessage;
use crate::server::WsServerAction;
use actix_web::{delete, get, post, put, web, HttpResponse};
use emgauwa_common::db::{DbController, DbJunctionRelaySchedule, DbSchedule, DbTag};
use emgauwa_common::errors::{ApiError, DatabaseError, EmgauwaError};
use emgauwa_common::models::{convert_db_list, FromDbModel, Schedule};
use emgauwa_common::types::{
ControllerWsAction, RequestScheduleCreate, RequestScheduleGetTagged, RequestScheduleUpdate,
ControllersWsAction, RequestScheduleCreate, RequestScheduleGetTagged, RequestScheduleUpdate,
ScheduleUid,
};
use itertools::Itertools;
use sqlx::pool::PoolConnection;
use sqlx::{Pool, Sqlite};
use crate::app_state;
use crate::app_state::AppState;
use crate::handlers::EmgauwaMessage;
use tokio::sync::mpsc;
#[get("/schedules")]
pub async fn index(pool: web::Data<Pool<Sqlite>>) -> Result<HttpResponse, EmgauwaError> {
@ -21,7 +19,7 @@ pub async fn index(pool: web::Data<Pool<Sqlite>>) -> Result<HttpResponse, Emgauw
let db_schedules = DbSchedule::get_all(&mut pool_conn).await?;
let schedules: Vec<Schedule> = convert_db_list(&mut pool_conn, db_schedules)?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().json(schedules))
}
@ -53,7 +51,7 @@ pub async fn tagged(
}
let schedules: Vec<Schedule> = convert_db_list(&mut pool_conn, db_schedules)?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().json(schedules))
}
@ -72,6 +70,7 @@ pub async fn show(
.ok_or(DatabaseError::NotFound)?;
let return_schedule = Schedule::from_db_model(&mut pool_conn, schedule)?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().json(return_schedule))
}
@ -97,6 +96,7 @@ pub async fn add(
}
let return_schedule = Schedule::from_db_model(&mut pool_conn, new_schedule)?;
pool_conn.close().await?;
Ok(HttpResponse::Created().json(return_schedule))
}
@ -133,13 +133,14 @@ pub async fn add_list(
}
let schedules: Vec<Schedule> = convert_db_list(&mut pool_conn, db_schedules)?;
pool_conn.close().await?;
Ok(HttpResponse::Created().json(schedules))
}
#[put("/schedules/{schedule_id}")]
pub async fn update(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
path: web::Path<(String,)>,
data: web::Json<RequestScheduleUpdate>,
) -> Result<HttpResponse, EmgauwaError> {
@ -179,15 +180,15 @@ pub async fn update(
let controller = DbController::get(&mut pool_conn, controller_id)
.await?
.ok_or(DatabaseError::NotFound)?;
app_state
.send(app_state::Action {
controller_uid: controller.uid,
action: ControllerWsAction::Schedules(vec![schedule.clone()]),
})
.await??;
ws_server_tx
.send(WsServerAction::ForwardToController(
controller.uid,
ControllersWsAction::Schedules(vec![schedule.clone()]),
))?;
}
let return_schedule = Schedule::from_db_model(&mut pool_conn, schedule)?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().json(return_schedule))
}
@ -196,8 +197,6 @@ pub async fn delete(
pool: web::Data<Pool<Sqlite>>,
path: web::Path<(String,)>,
) -> Result<HttpResponse, EmgauwaError> {
let mut pool_conn = pool.acquire().await?;
let (schedule_uid,) = path.into_inner();
let uid = ScheduleUid::try_from(schedule_uid.as_str())?;
@ -205,7 +204,9 @@ pub async fn delete(
ScheduleUid::Off => Err(EmgauwaError::from(ApiError::ProtectedSchedule)),
ScheduleUid::On => Err(EmgauwaError::from(ApiError::ProtectedSchedule)),
ScheduleUid::Any(_) => {
let mut pool_conn = pool.acquire().await?;
DbSchedule::delete_by_uid(&mut pool_conn, uid).await?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().emgauwa_message("schedule got deleted"))
}
}

View file

@ -12,6 +12,7 @@ pub async fn index(pool: web::Data<Pool<Sqlite>>) -> Result<HttpResponse, Emgauw
let mut pool_conn = pool.acquire().await?;
let db_tags = DbTag::get_all(&mut pool_conn).await?;
pool_conn.close().await?;
let tags: Vec<String> = db_tags.iter().map(|t| t.tag.clone()).collect();
@ -32,6 +33,7 @@ pub async fn show(
.ok_or(DatabaseError::NotFound)?;
let return_tag = Tag::from_db_model(&mut pool_conn, tag)?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().json(return_tag))
}
@ -45,6 +47,7 @@ pub async fn delete(
let (tag_name,) = path.into_inner();
DbTag::delete_by_tag(&mut pool_conn, &tag_name).await?;
pool_conn.close().await?;
Ok(HttpResponse::Ok().emgauwa_message("tag got deleted"))
}
@ -59,5 +62,6 @@ pub async fn add(
let cache = (Vec::new(), Vec::new()); // a new tag can't have any relays or schedules
let return_tag = Tag::from_db_model_cache(&mut pool_conn, new_tag, cache)?;
pool_conn.close().await?;
Ok(HttpResponse::Created().json(return_tag))
}

View file

@ -0,0 +1,114 @@
use std::time::Instant;
use actix_ws::{AggregatedMessage, CloseCode, CloseReason};
use futures::StreamExt;
use tokio::sync::{mpsc, oneshot};
use tokio::time::interval;
use emgauwa_common::constants::{HEARTBEAT_INTERVAL, HEARTBEAT_TIMEOUT};
use crate::server::{WsConnId, WsControllerAction, WsServerAction};
#[derive(Debug, Clone)]
pub struct ControllerWsHandle {
pub cmd_tx: mpsc::UnboundedSender<WsServerAction>,
}
impl ControllerWsHandle {
/// Register client message sender and obtain connection ID.
pub async fn connect(&self, conn_tx: mpsc::UnboundedSender<WsControllerAction>) -> WsConnId {
let (res_tx, res_rx) = oneshot::channel();
// unwrap: chat server should not have been dropped
self.cmd_tx
.send(WsServerAction::ConnectController(conn_tx, res_tx))
.unwrap();
// unwrap: chat server does not drop out the response channel
res_rx.await.unwrap()
}
/// Unregister message sender and broadcast disconnection message to current room.
pub fn disconnect(&self, conn: WsConnId) {
// unwrap: chat server should not have been dropped
self.cmd_tx.send(WsServerAction::DisconnectControllerConn(conn, true)).unwrap();
}
pub async fn message(&self, conn: WsConnId, msg: String) {
self.cmd_tx.send(WsServerAction::ControllerMessage(conn, msg)).unwrap();
}
pub async fn run(
&self,
mut session: actix_ws::Session,
msg_stream: actix_ws::MessageStream,
) {
log::info!("connected");
let mut last_heartbeat = Instant::now();
let mut interval = interval(HEARTBEAT_INTERVAL);
let (conn_tx, mut conn_rx) = mpsc::unbounded_channel();
// unwrap: chat server is not dropped before the HTTP server
let conn_id = self.connect(conn_tx).await;
let mut msg_stream = msg_stream
.max_frame_size(128 * 1024)
.aggregate_continuations()
.max_continuation_size(2 * 1024 * 1024);
let close_reason = loop {
tokio::select! {
Some(Ok(msg)) = msg_stream.next() => {
match msg {
AggregatedMessage::Ping(bytes) => {
last_heartbeat = Instant::now();
session.pong(&bytes).await.unwrap();
}
AggregatedMessage::Pong(_) => {
last_heartbeat = Instant::now();
}
AggregatedMessage::Text(text) => {
let text_str = text.to_string();
log::debug!("ws msg: {}", text_str);
self.message(conn_id, text_str).await;
}
AggregatedMessage::Binary(_) => {
log::warn!("unexpected binary message");
}
AggregatedMessage::Close(reason) => break reason,
}
}
Some(chat_msg) = conn_rx.recv() => {
match chat_msg {
WsControllerAction::Forward(msg) => {
session.text(msg).await.unwrap();
}
WsControllerAction::Disconnect => {
break Some(CloseReason::from((CloseCode::Normal, "Disconnect wanted by server")));
}
}
}
_ = interval.tick() => {
if Instant::now().duration_since(last_heartbeat) > HEARTBEAT_TIMEOUT {
break None;
}
let _ = session.ping(b"").await;
}
else => {
break None;
}
}
};
// attempt to close the connection gracefully
log::debug!("Closing a controller connection: {:?}", close_reason);
let _ = session.close(close_reason).await;
self.disconnect(conn_id);
}
}

View file

@ -1,114 +0,0 @@
use actix::{Actor, AsyncContext};
use emgauwa_common::db::{DbController, DbJunctionRelaySchedule, DbRelay, DbSchedule};
use emgauwa_common::errors::{DatabaseError, EmgauwaError};
use emgauwa_common::models::{Controller, FromDbModel};
use emgauwa_common::types::{ControllerWsAction, EmgauwaUid, RelayStates};
use emgauwa_common::utils;
use futures::executor::block_on;
use sqlx::pool::PoolConnection;
use sqlx::Sqlite;
use crate::app_state::{Action, ConnectController, UpdateRelayStates};
use crate::handlers::v1::ws::controllers::ControllersWs;
impl ControllersWs {
pub fn handle_register(
&mut self,
conn: &mut PoolConnection<Sqlite>,
ctx: &mut <ControllersWs as Actor>::Context,
controller: Controller,
) -> Result<(), EmgauwaError> {
log::info!(
"Registering controller: {} ({})",
controller.c.name,
controller.c.uid
);
let c = &controller.c;
let controller_db = block_on(DbController::get_by_uid_or_create(
conn,
&c.uid,
&c.name,
c.relay_count,
))?;
block_on(controller_db.update_active(conn, true))?;
// update only the relay count
block_on(controller_db.update(conn, &controller_db.name, c.relay_count))?;
for relay in &controller.relays {
log::debug!(
"Registering relay: {} ({})",
relay.r.name,
match relay.is_on {
Some(true) => "+",
Some(false) => "-",
None => "?",
}
);
let (new_relay, created) = block_on(DbRelay::get_by_controller_and_num_or_create(
conn,
&controller_db,
relay.r.number,
&relay.r.name,
))?;
if created {
let mut relay_schedules = Vec::new();
for schedule in &relay.schedules {
let (new_schedule, _) = block_on(DbSchedule::get_by_uid_or_create(
conn,
schedule.uid.clone(),
&schedule.name,
&schedule.periods,
))?;
relay_schedules.push(new_schedule);
}
block_on(DbJunctionRelaySchedule::set_schedules(
conn,
&new_relay,
relay_schedules.iter().collect(),
))?;
}
}
let controller_uid = &controller.c.uid;
let controller_db = block_on(DbController::get_by_uid(conn, controller_uid))?
.ok_or(DatabaseError::InsertGetError)?;
let controller = Controller::from_db_model(conn, controller_db)?;
let addr = ctx.address();
self.controller_uid = Some(controller_uid.clone());
block_on(self.app_state.send(ConnectController {
address: addr.recipient(),
controller: controller.clone(),
}))??;
block_on(self.app_state.send(Action {
controller_uid: controller_uid.clone(),
action: ControllerWsAction::Controller(controller.clone()),
}))??;
block_on(self.app_state.send(Action {
controller_uid: controller_uid.clone(),
action: ControllerWsAction::Relays(controller.relays),
}))??;
log::debug!("Done registering controller");
Ok(())
}
pub fn handle_relay_states(
&mut self,
controller_uid: EmgauwaUid,
relay_states: RelayStates,
) -> Result<(), EmgauwaError> {
log::debug!(
"Received relay states: {} for {}",
utils::printable_relay_states(&relay_states),
controller_uid
);
block_on(self.app_state.send(UpdateRelayStates {
controller_uid,
relay_states,
}))?;
Ok(())
}
}

View file

@ -1,153 +1,53 @@
use std::time::Instant;
use actix::{Actor, ActorContext, Addr, AsyncContext, Handler, StreamHandler};
use actix_web_actors::ws;
use actix_web_actors::ws::ProtocolError;
use futures::executor::block_on;
use sqlx::{Pool, Sqlite};
use sqlx::pool::PoolConnection;
use ws::Message;
use emgauwa_common::constants::{HEARTBEAT_INTERVAL, HEARTBEAT_TIMEOUT};
use actix_web::{get, web, HttpRequest, HttpResponse};
use tokio::sync::mpsc;
use tokio::task::spawn_local;
use emgauwa_common::constants;
use emgauwa_common::errors::EmgauwaError;
use emgauwa_common::types::{ControllerWsAction, EmgauwaUid};
use crate::handlers::v1::ws::controllers::handle::ControllerWsHandle;
use crate::server::WsServerAction;
use crate::settings::Settings;
use crate::app_state::{AppState, DisconnectController};
mod handle;
mod handlers;
pub struct ControllersWs {
pub pool: Pool<Sqlite>,
pub controller_uid: Option<EmgauwaUid>,
pub app_state: Addr<AppState>,
pub hb: Instant,
pub async fn run_ws_server_handle(ws_server_tx: mpsc::UnboundedSender<WsServerAction>, session: actix_ws::Session, msg_stream: actix_ws::MessageStream) {
let ws_server_handle = ControllerWsHandle {
cmd_tx: ws_server_tx,
};
ws_server_handle.run(session, msg_stream).await;
}
impl Actor for ControllersWs {
type Context = ws::WebsocketContext<Self>;
//noinspection DuplicatedCode
#[get("/ws/controllers")]
pub async fn ws_controllers(
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
settings: web::Data<Settings>,
req: HttpRequest,
stream: web::Payload,
) -> Result<HttpResponse, EmgauwaError> {
let token = req
.headers()
.get(constants::CONTROLLER_WS_TOKEN_HEADER)
.ok_or(EmgauwaError::Unauthorized(
String::from("Missing or invalid token header"),
))?
.to_str()
.map_err(|_| EmgauwaError::Unauthorized(String::from("Invalid token header")))?;
fn started(&mut self, ctx: &mut Self::Context) {
self.hb(ctx);
}
if token.ne(&settings.server.token) {
return Err(EmgauwaError::Unauthorized(String::from("Wrong token header")));
}
fn stopped(&mut self, _ctx: &mut Self::Context) {
if let Some(controller_uid) = &self.controller_uid {
if let Err(err) = block_on(self.app_state.send(DisconnectController {
controller_uid: controller_uid.clone(),
})).unwrap_or_else(|err| Err(EmgauwaError::from(err))) {
log::error!("Error disconnecting controller: {:?}", err);
}
}
}
}
impl ControllersWs {
pub fn handle_action(
&mut self,
conn: &mut PoolConnection<Sqlite>,
ctx: &mut <ControllersWs as Actor>::Context,
action: ControllerWsAction,
) {
let action_res = match action {
ControllerWsAction::Register(controller) => self.handle_register(conn, ctx, controller),
ControllerWsAction::RelayStates((controller_uid, relay_states)) => {
self.handle_relay_states(controller_uid, relay_states)
}
_ => Ok(()),
};
if let Err(e) = action_res {
log::error!("Error handling action: {:?}", e);
ctx.text(
serde_json::to_string(&e).unwrap_or(format!("Error in handling action: {:?}", e)),
);
}
}
// helper method that sends ping to client every 5 seconds (HEARTBEAT_INTERVAL).
fn hb(&self, ctx: &mut ws::WebsocketContext<Self>) {
ctx.run_interval(HEARTBEAT_INTERVAL, |act, ctx| {
// check client heartbeats
if Instant::now().duration_since(act.hb) > HEARTBEAT_TIMEOUT {
log::warn!("Websocket Controller heartbeat failed, disconnecting!");
ctx.stop();
// don't try to send a ping
return;
}
log::trace!("Sending ping to controller");
ctx.ping(&[]);
});
}
}
impl Handler<ControllerWsAction> for ControllersWs {
type Result = Result<(), EmgauwaError>;
fn handle(&mut self, action: ControllerWsAction, ctx: &mut Self::Context) -> Self::Result {
match action {
ControllerWsAction::Disconnect => {
ctx.close(None);
ctx.stop();
}
_ => {
let action_json = serde_json::to_string(&action)?;
ctx.text(action_json);
}
}
Ok(())
}
}
impl StreamHandler<Result<Message, ProtocolError>> for ControllersWs {
fn handle(&mut self, msg: Result<Message, ProtocolError>, ctx: &mut Self::Context) {
let mut pool_conn = match block_on(self.pool.acquire()) {
Ok(conn) => conn,
Err(err) => {
log::error!("Failed to acquire database connection: {:?}", err);
ctx.stop();
return;
}
};
let msg = match msg {
Err(_) => {
ctx.stop();
return;
}
Ok(msg) => msg,
};
match msg {
Message::Ping(msg) => {
log::trace!("Received ping from controller: {:?}", msg);
self.hb = Instant::now();
ctx.pong(&msg)
}
Message::Pong(_) => {
log::trace!("Received pong from controller");
self.hb = Instant::now();
}
Message::Text(text) => match serde_json::from_str(&text) {
Ok(action) => {
self.handle_action(&mut pool_conn, ctx, action);
}
Err(e) => {
log::error!("Error deserializing action: {:?}", e);
ctx.text(
serde_json::to_string(&EmgauwaError::Serialization(e))
.unwrap_or(String::from("Error in deserializing action")),
);
}
},
Message::Binary(_) => log::warn!("Received unexpected binary in controller ws"),
Message::Close(reason) => {
ctx.close(reason);
ctx.stop();
}
Message::Continuation(_) => {
ctx.stop();
}
Message::Nop => (),
}
}
let (res, session, msg_stream) = actix_ws::handle(&req, stream).map_err(|e| {
log::error!("Error in websocket handshake: {}", e);
EmgauwaError::Internal(String::from("Error in websocket handshake"))
})?;
// spawn websocket handler (and don't await it) so that the response is returned immediately
spawn_local(run_ws_server_handle(
(**ws_server_tx).clone(),
session,
msg_stream,
));
Ok(res)
}

View file

@ -1,68 +1,4 @@
use std::time::Instant;
use actix::Addr;
use actix_web::{get, web, HttpRequest, HttpResponse};
use actix_web_actors::ws;
use emgauwa_common::errors::EmgauwaError;
use sqlx::{Pool, Sqlite};
use emgauwa_common::constants;
use crate::app_state::AppState;
use crate::handlers::v1::ws::controllers::ControllersWs;
use crate::handlers::v1::ws::relays::RelaysWs;
use crate::settings::Settings;
pub mod controllers;
pub mod relays;
#[get("/ws/controllers")]
pub async fn ws_controllers(
pool: web::Data<Pool<Sqlite>>,
app_state: web::Data<Addr<AppState>>,
settings: web::Data<Settings>,
req: HttpRequest,
stream: web::Payload,
) -> Result<HttpResponse, EmgauwaError> {
let token = req
.headers()
.get(constants::CONTROLLER_WS_TOKEN_HEADER)
.ok_or(EmgauwaError::Unauthorized(
String::from("Missing or invalid token header"),
))?
.to_str()
.map_err(|_| EmgauwaError::Unauthorized(String::from("Invalid token header")))?;
if token != settings.server.token {
return Err(EmgauwaError::Unauthorized(String::from("Wrong token header")));
}
let resp = ws::start(
ControllersWs {
pool: pool.get_ref().clone(),
controller_uid: None,
app_state: app_state.get_ref().clone(),
hb: Instant::now(),
},
&req,
stream,
)
.map_err(|_| EmgauwaError::Internal(String::from("error starting websocket")));
resp
}
#[get("/ws/relays")]
pub async fn ws_relays(
app_state: web::Data<Addr<AppState>>,
req: HttpRequest,
stream: web::Payload,
) -> Result<HttpResponse, EmgauwaError> {
let resp = ws::start(
RelaysWs {
app_state: app_state.get_ref().clone(),
hb: Instant::now(),
},
&req,
stream,
)
.map_err(|_| EmgauwaError::Internal(String::from("error starting websocket")));
resp
}
mod relays;
pub use relays::ws_relays;
mod controllers;
pub use controllers::ws_controllers;

View file

@ -0,0 +1,108 @@
use std::time::Instant;
use actix_ws::{AggregatedMessage, CloseCode, CloseReason};
use futures::StreamExt;
use tokio::sync::{mpsc, oneshot};
use tokio::time::interval;
use emgauwa_common::constants::{HEARTBEAT_INTERVAL, HEARTBEAT_TIMEOUT};
use crate::server::{WsConnId, WsRelayAction, WsServerAction};
#[derive(Debug, Clone)]
pub struct RelayWsHandle {
pub cmd_tx: mpsc::UnboundedSender<WsServerAction>,
}
impl RelayWsHandle {
/// Register client message sender and obtain connection ID.
pub async fn connect(&self, conn_tx: mpsc::UnboundedSender<WsRelayAction>) -> WsConnId {
let (res_tx, res_rx) = oneshot::channel();
// unwrap: chat server should not have been dropped
self.cmd_tx
.send(WsServerAction::ConnectRelay(conn_tx, res_tx))
.unwrap();
// unwrap: chat server does not drop out the response channel
res_rx.await.unwrap()
}
/// Unregister message sender and broadcast disconnection message to current room.
pub fn disconnect(&self, conn: WsConnId) {
// unwrap: chat server should not have been dropped
self.cmd_tx.send(WsServerAction::DisconnectRelayConn(conn)).unwrap();
}
pub async fn run(
&self,
mut session: actix_ws::Session,
msg_stream: actix_ws::MessageStream,
) {
log::info!("connected");
let mut last_heartbeat = Instant::now();
let mut interval = interval(HEARTBEAT_INTERVAL);
let (conn_tx, mut conn_rx) = mpsc::unbounded_channel();
// unwrap: chat server is not dropped before the HTTP server
let conn_id = self.connect(conn_tx).await;
let mut msg_stream = msg_stream
.max_frame_size(128 * 1024)
.aggregate_continuations()
.max_continuation_size(2 * 1024 * 1024);
let close_reason = loop {
tokio::select! {
Some(Ok(msg)) = msg_stream.next() => {
match msg {
AggregatedMessage::Ping(bytes) => {
last_heartbeat = Instant::now();
session.pong(&bytes).await.unwrap();
}
AggregatedMessage::Pong(_) => {
last_heartbeat = Instant::now();
}
AggregatedMessage::Text(_) => {
log::warn!("unexpected text message");
}
AggregatedMessage::Binary(_) => {
log::warn!("unexpected binary message");
}
AggregatedMessage::Close(reason) => break reason,
}
}
Some(chat_msg) = conn_rx.recv() => {
match chat_msg {
WsRelayAction::Forward(msg) => {
session.text(msg).await.unwrap();
}
WsRelayAction::Disconnect => {
break Some(CloseReason::from((CloseCode::Normal, "Disconnect wanted by server")));
}
}
}
_ = interval.tick() => {
if Instant::now().duration_since(last_heartbeat) > HEARTBEAT_TIMEOUT {
break None;
}
let _ = session.ping(b"").await;
}
else => {
break None;
}
}
};
self.disconnect(conn_id);
// attempt to close the connection gracefully
let _ = session.close(close_reason).await;
}
}

View file

@ -1,109 +1,36 @@
use std::time::Instant;
use actix::{Actor, ActorContext, Addr, AsyncContext, Handler, Message, StreamHandler};
use actix_web_actors::ws;
use actix_web_actors::ws::ProtocolError;
use emgauwa_common::constants::{HEARTBEAT_INTERVAL, HEARTBEAT_TIMEOUT};
use actix_web::{get, web, HttpRequest, HttpResponse};
use emgauwa_common::errors::EmgauwaError;
use futures::executor::block_on;
use tokio::sync::mpsc;
use tokio::task::spawn_local;
use crate::server::WsServerAction;
use crate::app_state::{AppState, ConnectRelayClient};
mod handle;
pub struct RelaysWs {
pub app_state: Addr<AppState>,
pub hb: Instant,
pub async fn run_ws_server_handle(ws_server_tx: mpsc::UnboundedSender<WsServerAction>, session: actix_ws::Session, msg_stream: actix_ws::MessageStream) {
let ws_server_handle = crate::handlers::v1::ws::relays::handle::RelayWsHandle {
cmd_tx: ws_server_tx,
};
ws_server_handle.run(session, msg_stream).await;
}
#[derive(Message)]
#[rtype(result = "()")]
pub struct SendRelays {
pub relays_json: String,
}
impl Actor for RelaysWs {
type Context = ws::WebsocketContext<Self>;
fn started(&mut self, ctx: &mut Self::Context) {
// get unique id for ctx
match self.get_relays_json() {
Ok(relays_json) => {
ctx.text(relays_json);
self.hb(ctx);
block_on(self.app_state.send(ConnectRelayClient {
addr: ctx.address(),
})).map_err(
|err| log::error!("Error connecting relay-client: {:?}", err)
).ok();
}
Err(err) => {
log::error!("Error getting relays: {:?}", err);
ctx.stop();
}
}
}
}
impl RelaysWs {
fn get_relays_json(&self) -> Result<String, EmgauwaError> {
let relays = block_on(self.app_state.send(crate::app_state::GetRelays {}))??;
serde_json::to_string(&relays).map_err(EmgauwaError::from)
}
// helper method that sends ping to client every 5 seconds (HEARTBEAT_INTERVAL).
fn hb(&self, ctx: &mut ws::WebsocketContext<Self>) {
ctx.run_interval(HEARTBEAT_INTERVAL, |act, ctx| {
// check client heartbeats
if Instant::now().duration_since(act.hb) > HEARTBEAT_TIMEOUT {
log::debug!("Websocket Relay heartbeat failed, disconnecting!");
ctx.stop();
// don't try to send a ping
return;
}
ctx.ping(&[]);
});
}
}
impl StreamHandler<Result<ws::Message, ProtocolError>> for RelaysWs {
fn handle(&mut self, msg: Result<ws::Message, ProtocolError>, ctx: &mut Self::Context) {
let msg = match msg {
Err(_) => {
ctx.stop();
return;
}
Ok(msg) => msg,
};
match msg {
ws::Message::Ping(msg) => {
log::trace!("Received ping from relay-client: {:?}", msg);
self.hb = Instant::now();
ctx.pong(&msg)
}
ws::Message::Pong(_) => {
log::trace!("Received pong from relay-client");
self.hb = Instant::now();
}
ws::Message::Text(_) => log::debug!("Received unexpected text in relays ws"),
ws::Message::Binary(_) => log::debug!("Received unexpected binary in relays ws"),
ws::Message::Close(reason) => {
ctx.close(reason);
ctx.stop();
}
ws::Message::Continuation(_) => {
ctx.stop();
}
ws::Message::Nop => (),
}
}
}
impl Handler<SendRelays> for RelaysWs {
type Result = ();
fn handle(&mut self, msg: SendRelays, ctx: &mut Self::Context) -> Self::Result {
ctx.text(msg.relays_json);
}
//noinspection DuplicatedCode
#[get("/ws/relays")]
pub async fn ws_relays(
ws_server_tx: web::Data<mpsc::UnboundedSender<WsServerAction>>,
req: HttpRequest,
stream: web::Payload,
) -> Result<HttpResponse, EmgauwaError> {
let (res, session, msg_stream) = actix_ws::handle(&req, stream).map_err(|e| {
log::error!("Error in websocket handshake: {}", e);
EmgauwaError::Internal(String::from("Error in websocket handshake"))
})?;
// spawn websocket handler (and don't await it) so that the response is returned immediately
spawn_local(run_ws_server_handle(
(**ws_server_tx).clone(),
session,
msg_stream,
));
Ok(res)
}