use std::{collections::HashMap, sync::atomic::AtomicU32};
use crate::{NodeId, PeerId};
use ruda_communication::{Address, CommunicationError};
use ruda_core::id::IdGenerator;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Hash)]
pub(crate) struct RequestId(u32);
static REQ_ID_COUNTER: AtomicU32 = AtomicU32::new(0);
impl RequestId {
pub(crate) fn new() -> Self {
let id = REQ_ID_COUNTER.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Self(id)
}
}
impl Default for RequestId {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, PartialEq, Eq, Clone, Copy, Hash, Serialize, Deserialize, PartialOrd, Ord)]
pub(crate) struct SessionId {
id: u64,
}
impl SessionId {
pub(crate) fn new() -> Self {
Self {
id: IdGenerator::generate(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) enum CollectiveMessage {
Init(SessionId),
Request(RequestId, RemoteRequest),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) struct CollectiveMessageResponse {
pub request_id: RequestId,
pub content: RemoteResponse,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) enum RemoteRequest {
Register {
node_addr: Address,
num_nodes: u32,
peers: Vec<PeerId>,
},
Finish,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub(crate) enum RemoteResponse {
Register {
node_id: NodeId,
nodes: HashMap<NodeId, Address>,
num_global_devices: u32,
},
FinishAck,
Error(GlobalCollectiveError),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum GlobalCollectiveError {
AllReduceBeforeRegister,
RingReduceImpossible,
NotRegisteredOnFinish,
PendingRegisterOnFinish,
RegisterParamsMismatch,
DoubleRegister,
AllReduceParamsMismatch,
FirstMsgNotInit,
InvalidMessage,
PeerSentIncoherentTensor,
PeerLost(NodeId),
Server(String),
WrongOrchestratorResponse,
OrchestratorUnreachable,
}
impl<E: CommunicationError> From<E> for GlobalCollectiveError {
fn from(err: E) -> Self {
Self::Server(format!("{err:?}"))
}
}