use std::fmt::Debug;
use std::sync::Arc;
use tokio::sync::Mutex;
use crate::global::{
orchestrator::state::GlobalCollectiveState,
shared::{CollectiveMessage, GlobalCollectiveError},
};
use ruda_communication::{
CommunicationChannel, Message, ProtocolServer, util::os_shutdown_signal, websocket::WsServer,
};
#[derive(Clone)]
pub(crate) struct GlobalOrchestrator {
state: Arc<Mutex<GlobalCollectiveState>>,
}
impl GlobalOrchestrator {
pub(crate) async fn start<F, S: ProtocolServer + Debug>(
shutdown_signal: F,
comms_server: S,
) -> Result<(), GlobalCollectiveError>
where
F: Future<Output = ()> + Send + 'static,
{
let state = GlobalCollectiveState::new();
let server = Self {
state: Arc::new(tokio::sync::Mutex::new(state)),
};
comms_server
.route("/response", {
let server = server.clone();
async move |socket| {
if let Err(err) = server.handle_socket_response::<S>(socket).await {
log::error!("[Response Handler] Error: {err:?}")
}
}
})
.route("/request", {
let server = server.clone();
async move |socket| {
if let Err(err) = server.handle_socket_request::<S>(socket).await {
log::error!("[Request Handler] Error: {err:?}")
}
}
})
.serve(shutdown_signal)
.await
.map_err(|err| GlobalCollectiveError::Server(format!("{err:?}")))?;
Ok(())
}
async fn handle_socket_response<S: ProtocolServer>(
self,
mut stream: S::Channel,
) -> Result<(), GlobalCollectiveError> {
log::info!("[Response Handler] On new connection.");
let policy = crate::global::policy::GlobalFailurePolicy::from_environment().map_err(GlobalCollectiveError::Server)?;
let msg = tokio::time::timeout(policy.connect_timeout, stream.recv()).await
.map_err(|_| GlobalCollectiveError::OperationTimeout)??;
let Some(msg) = msg else {
log::warn!("Response socket closed early!");
return Ok(());
};
let msg = rmp_serde::from_slice::<CollectiveMessage>(&msg.data)
.map_err(|_| GlobalCollectiveError::InvalidMessage)?;
let CollectiveMessage::Init(id) = msg else {
return Err(GlobalCollectiveError::FirstMsgNotInit);
};
let mut receiver = {
let mut state = self.state.lock().await;
state.get_session_responder(id)?
};
let result = async {
loop {
tokio::select! {
packet = stream.recv() => {
match packet? {
None => break,
Some(_) => return Err(GlobalCollectiveError::InvalidMessage),
}
}
response = receiver.recv() => {
let Some(response) = response else { break; };
let bytes = rmp_serde::to_vec(&response).map_err(|_| GlobalCollectiveError::InvalidMessage)?;
tokio::time::timeout(policy.request_timeout, stream.send(Message::new(bytes.into())))
.await.map_err(|_| GlobalCollectiveError::OperationTimeout)??;
}
}
}
Ok(())
}.await;
self.state.lock().await.disconnect(id).await;
result
}
async fn handle_socket_request<S: ProtocolServer>(
self,
mut stream: S::Channel,
) -> Result<(), GlobalCollectiveError> {
log::info!("[Request Handler] On new connection.");
let mut session_id = None;
let policy = crate::global::policy::GlobalFailurePolicy::from_environment().map_err(GlobalCollectiveError::Server)?;
let result = async {
loop {
let packet = if session_id.is_none() {
tokio::time::timeout(policy.connect_timeout, stream.recv()).await
.map_err(|_| GlobalCollectiveError::OperationTimeout)??
} else { stream.recv().await? };
let Some(msg) = packet else {
log::info!("Peer closed the connection");
break;
};
let mut state = self.state.lock().await;
let msg = rmp_serde::from_slice::<CollectiveMessage>(&msg.data)
.map_err(|_| GlobalCollectiveError::InvalidMessage)?;
match msg {
CollectiveMessage::Init(id) => {
if session_id.is_some() { return Err(GlobalCollectiveError::InvalidMessage); }
state.init_session(id);
session_id = Some(id);
}
CollectiveMessage::Request(request_id, remote_request) => {
let session_id = session_id.ok_or(GlobalCollectiveError::FirstMsgNotInit)?;
state
.process_request(session_id, request_id, remote_request)
.await;
}
}
}
Ok(())
}.await;
if let Some(id) = session_id { self.state.lock().await.disconnect(id).await; }
result
}
}
pub async fn start_global_orchestrator(port: u16) {
let server = WsServer::new(port);
let res = GlobalOrchestrator::start(os_shutdown_signal(), server).await;
if let Err(err) = res {
log::error!("Global Collective Orchestrator error: {err:?}");
}
}