use crate::app::AppState; use crate::context::RequestContext; use crate::error::WebError; 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; use serde_json::{json, Value}; pub async fn socket( ws: WebSocketUpgrade, State(state): State, Extension(context): Extension, Query(query): Query, ) -> Result { 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 _ = socket.send(serialize_snapshot_message(&payload)).await; while let Some(message) = socket.next().await { let Ok(message) = message else { break; }; match message { 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; } } Err(error) => { let 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(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; } } } Message::Close(_) => 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::(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\"")); } }