use aion_core::{ClusterDeployment, ClusterSnapshot, ClusterStreamError, ClusterWorker};
use aion_proto::{StreamedClusterEvent, StreamedClusterSnapshot, WireError};
use axum::extract::ws::{CloseFrame, Message, WebSocket, close_code};
use futures::{SinkExt, StreamExt};
use crate::cluster_publisher::ClusterStreamLagged;
use crate::error::ServerError;
use crate::namespace::CallerIdentity;
use crate::state::ServerState;
pub async fn serve_cluster_socket(
mut socket: WebSocket,
state: &ServerState,
caller: &CallerIdentity,
after_seq: u64,
) -> Result<(), ServerError> {
if !caller.deploy_granted() {
let error = ServerError::namespace_denied(
"cluster topology subscription requires the deployment-wide deploy grant",
);
super::socket::send_wire_error(&mut socket, &error.to_wire_error()).await?;
return Err(error);
}
let publisher = state.cluster_publisher();
let mut live = publisher.subscribe(after_seq);
let snapshot = build_snapshot(state, caller).await?;
let priming = StreamedClusterSnapshot::new(snapshot);
let priming = serde_json::to_string(&priming).map_err(|source| ServerError::Wire {
wire: WireError::backend(format!(
"failed to serialize cluster snapshot frame: {source}"
)),
})?;
if socket.send(Message::Text(priming.into())).await.is_err() {
return Ok(());
}
let (mut socket_tx, mut socket_rx) = socket.split();
loop {
tokio::select! {
client_message = socket_rx.next() => {
match client_message {
Some(Ok(Message::Close(_))) | None => return send_normal_close(&mut socket_tx).await,
Some(Ok(_other)) => {}
Some(Err(_error)) => return Ok(()),
}
}
item = live.next() => {
match item {
Some(Ok(event)) => {
let frame = StreamedClusterEvent::new(event);
let frame = match serde_json::to_string(&frame) {
Ok(frame) => frame,
Err(source) => {
let error = ServerError::Wire {
wire: WireError::backend(format!(
"failed to serialize cluster event frame: {source}"
)),
};
super::socket::send_wire_error(&mut socket_tx, &error.to_wire_error()).await?;
return Err(error);
}
};
if socket_tx.send(Message::Text(frame.into())).await.is_err() {
return Ok(());
}
}
Some(Err(ClusterStreamLagged { skipped })) => {
let lagged = ClusterStreamError::ClusterLagged { skipped };
return deliver_cluster_terminal(&mut socket_tx, &lagged).await;
}
None => {
return send_normal_close(&mut socket_tx).await;
}
}
}
}
}
}
pub(crate) async fn build_snapshot(
state: &ServerState,
caller: &CallerIdentity,
) -> Result<ClusterSnapshot, ServerError> {
if !caller.deploy_granted() {
return Err(ServerError::namespace_denied(
"cluster topology snapshot requires the deployment-wide deploy grant",
));
}
let node = state.cluster_self_node().map_or_else(
|| STANDALONE_NODE_LABEL.to_owned(),
std::borrow::ToOwned::to_owned,
);
let workers = state
.worker_registry()
.all_workers()?
.into_iter()
.map(|handle| ClusterWorker {
deployment: handle
.instance()
.map(|instance| instance.deployment.clone()),
worker_id: handle.id().value().to_string(),
namespaces: handle.namespaces().iter().cloned().collect(),
task_queue: handle.task_queue().to_owned(),
transport: handle.delivery().transport(),
node: handle.node().map(str::to_owned),
deployment_association: handle.instance().map(|instance| instance.association),
})
.collect();
let listing = state
.worker_deployment_store()
.list_worker_deployments()
.await?;
for row in listing.undecodable {
tracing::warn!(
name = %row.name,
error = %row.error,
"worker deployment row could not be decoded while building cluster snapshot"
);
}
let deployments = listing
.deployments
.into_iter()
.map(|record| ClusterDeployment {
name: record.name,
desired_state: record.desired,
binary_version: record.binary.version,
binary_content_hash: record.binary.content_hash,
})
.collect();
Ok(ClusterSnapshot {
node,
as_of_seq: state.cluster_publisher().current_seq(),
peers: Vec::new(),
shards: Vec::new(),
workers,
deployments,
})
}
const STANDALONE_NODE_LABEL: &str = "standalone";
async fn deliver_cluster_terminal<Tx>(
socket_tx: &mut Tx,
error: &ClusterStreamError,
) -> Result<(), ServerError>
where
Tx: futures::Sink<Message> + Unpin,
<Tx as futures::Sink<Message>>::Error: std::fmt::Debug,
{
let payload = serde_json::json!({ "error": error });
let payload = serde_json::to_string(&payload).map_err(|source| ServerError::Wire {
wire: WireError::backend(format!(
"failed to serialize cluster stream error: {source}"
)),
})?;
if socket_tx.send(Message::Text(payload.into())).await.is_ok() {
let close = CloseFrame {
code: close_code::ERROR,
reason: "cluster_lagged".into(),
};
let close_result = socket_tx.send(Message::Close(Some(close))).await;
drop(close_result);
}
Err(ServerError::lagged_stream())
}
async fn send_normal_close<Tx>(socket_tx: &mut Tx) -> Result<(), ServerError>
where
Tx: futures::Sink<Message> + Unpin,
<Tx as futures::Sink<Message>>::Error: std::fmt::Debug,
{
let close = CloseFrame {
code: close_code::NORMAL,
reason: "subscription complete".into(),
};
let close_result = socket_tx.send(Message::Close(Some(close))).await;
drop(close_result);
Ok(())
}