Files
gree-controller/src/api/websocket.rs
T

56 lines
1.9 KiB
Rust

#[derive(Debug, Deserialize)]
struct WsQuery {
token: Option<String>,
}
async fn websocket(
State(state): State<AppState>,
Query(query): Query<WsQuery>,
ws: WebSocketUpgrade,
) -> Result<Response, AppError> {
let expected = state.config.app_token.trim();
if !expected.is_empty() && query.token.as_deref() != Some(expected) {
return Err(AppError::Unauthorized);
}
Ok(ws.on_upgrade(move |socket| websocket_loop(state, socket)))
}
async fn websocket_loop(state: AppState, mut socket: WebSocket) {
let initial = match build_bootstrap(&state).await {
Ok(data) => json!({"event":"bootstrap","timestamp":Utc::now(),"data":data}),
Err(err) => {
json!({"event":"error","timestamp":Utc::now(),"data":{"message":err.to_string()}})
}
};
if socket
.send(Message::Text(initial.to_string()))
.await
.is_err()
{
return;
}
let mut receiver = state.events.subscribe();
loop {
tokio::select! {
event = receiver.recv() => {
match event {
Ok(event) => {
if let Ok(text) = serde_json::to_string(&event) {
if socket.send(Message::Text(text)).await.is_err() { break; }
}
}
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => continue,
Err(_) => break,
}
}
message = socket.next() => {
match message {
Some(Ok(Message::Ping(value))) => { if socket.send(Message::Pong(value)).await.is_err() { break; } }
Some(Ok(Message::Text(text))) if text == "ping" => { if socket.send(Message::Text("pong".into())).await.is_err() { break; } }
Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
_ => {}
}
}
}
}
}