Files
mnote/rust/crates/mnote-web/src/routes/ws.rs
T

187 lines
7.2 KiB
Rust
Raw Normal View History

use crate::app::AppState;
use crate::context::RequestContext;
use crate::error::WebError;
2026-04-29 12:24:44 +08:00
use crate::routes::stream_support::{load_stream_snapshot, StreamSnapshotQuery};
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::extract::{Extension, Query, State};
use axum::response::Response;
use futures_util::StreamExt;
2026-04-29 12:24:44 +08:00
use serde_json::{json, Value};
use tokio::sync::broadcast::error::RecvError;
pub async fn socket(
ws: WebSocketUpgrade,
State(state): State<AppState>,
Extension(context): Extension<RequestContext>,
Query(query): Query<StreamSnapshotQuery>,
) -> Result<Response, WebError> {
let snapshot = load_stream_snapshot(state.config(), &context, &query).await?;
let state = state.clone();
let context = context.clone();
let query = query.clone();
Ok(ws.on_upgrade(move |socket| handle_socket(socket, state, context, query, snapshot)))
}
async fn handle_socket(
mut socket: WebSocket,
state: AppState,
context: RequestContext,
query: StreamSnapshotQuery,
payload: Value,
) {
let mut stream_delta_rx = state.stream_delta_tx.subscribe();
let _ = socket.send(serialize_snapshot_message(&payload)).await;
loop {
tokio::select! {
biased;
// 优先处理 broadcast 推送的变更通知
delta_result = stream_delta_rx.recv() => {
match delta_result {
Ok(delta) => {
let notify = json!({
"kind": "delta",
"data": delta,
"requestId": context.trace.request_id,
"traceId": context.trace.trace_id,
"workspaceId": delta.get("workspaceId").and_then(Value::as_str).unwrap_or(""),
});
if socket.send(Message::Text(notify.to_string().into())).await.is_err() {
break;
}
}
Err(RecvError::Lagged(n)) => {
// Lagged: 发送 resync 提示让客户端重新加载
let lagged_hint = json!({
"kind": "resync_hint",
"reason": "stream lagged",
"dropped": n,
"requestId": context.trace.request_id,
"traceId": context.trace.trace_id,
});
let _ = socket.send(Message::Text(lagged_hint.to_string().into())).await;
}
Err(RecvError::Closed) => {
// Broadcast channel closed, WS stays open for client-initiated resync
}
}
}
// 处理客户端消息(resync 请求等)
message = socket.next() => {
match message {
Some(Ok(Message::Text(text))) => {
if is_resync_request(&text) {
match load_stream_snapshot(state.config(), &context, &query).await {
Ok(snapshot) => {
if socket
.send(serialize_resync_message(&snapshot))
.await
.is_err()
{
break;
}
// 重订阅 broadcast(可能丢掉了中间的变更)
stream_delta_rx = state.stream_delta_tx.subscribe();
}
Err(error) => {
let err_payload = json!({
"kind": "error",
"code": "snapshot_reload_failed",
"message": format!("{error:?}"),
"requestId": context.trace.request_id,
"traceId": context.trace.trace_id,
});
if socket
.send(Message::Text(err_payload.to_string().into()))
.await
.is_err()
{
break;
}
}
}
} else {
let ack = json!({
"kind": "ack",
"requestId": context.trace.request_id,
"traceId": context.trace.trace_id,
"accepted": false,
"reason": "unsupported_message",
});
if socket.send(Message::Text(ack.to_string().into())).await.is_err() {
break;
}
}
}
Some(Ok(Message::Close(_))) => break,
Some(Ok(_)) => {} // ignore binary/ping/pong
Some(Err(_)) => break,
None => break,
}
}
}
}
}
fn serialize_snapshot_message(payload: &Value) -> Message {
Message::Text(payload.to_string().into())
}
fn serialize_resync_message(payload: &Value) -> Message {
Message::Text(
json!({
"kind": "resync",
"snapshot": payload,
})
.to_string()
.into(),
)
}
fn is_resync_request(text: &str) -> bool {
let trimmed = text.trim();
if trimmed.eq_ignore_ascii_case("resync") {
return true;
}
serde_json::from_str::<Value>(trimmed)
.ok()
.and_then(|value| value.get("type").and_then(Value::as_str).map(str::to_owned))
.map(|value| value.eq_ignore_ascii_case("resync"))
.unwrap_or(false)
}
#[cfg(test)]
mod tests {
use super::{is_resync_request, serialize_resync_message, serialize_snapshot_message};
use axum::extract::ws::Message;
use serde_json::json;
#[test]
fn ws_resync_detection_accepts_plain_text_and_json() {
assert!(is_resync_request("resync"));
assert!(is_resync_request(r#"{"type":"resync"}"#));
assert!(!is_resync_request("hello"));
}
#[test]
fn ws_snapshot_serializers_emit_text_frames() {
let snapshot = json!({
"kind": "snapshot",
"scope": "workspace",
});
let Message::Text(snapshot_text) = serialize_snapshot_message(&snapshot) else {
panic!("snapshot message 应该是文本帧");
};
assert!(snapshot_text.contains("\"kind\":\"snapshot\""));
let Message::Text(resync_text) = serialize_resync_message(&snapshot) else {
panic!("resync message 应该是文本帧");
};
assert!(resync_text.contains("\"kind\":\"resync\""));
}
}