feat: 收口 tree-first graph 主链与前端测试修复
This commit is contained in:
@@ -1,30 +1,36 @@
|
||||
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;
|
||||
use axum::extract::{Extension, Query, State};
|
||||
use axum::response::Response;
|
||||
use futures_util::StreamExt;
|
||||
use serde_json::json;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
pub async fn socket(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<AppState>,
|
||||
Extension(context): Extension<RequestContext>,
|
||||
) -> Response {
|
||||
ws.on_upgrade(move |socket| handle_socket(socket, context))
|
||||
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, context: RequestContext) {
|
||||
let payload = json!({
|
||||
"kind": "ws_placeholder",
|
||||
"requestId": context.trace.request_id,
|
||||
"traceId": context.trace.trace_id,
|
||||
"workspaceId": context.workspace.workspace_id,
|
||||
"notes": [
|
||||
"当前为 task-062 最小骨架,后续可承接协作推送和运行态事件。"
|
||||
]
|
||||
});
|
||||
|
||||
async fn handle_socket(
|
||||
mut socket: WebSocket,
|
||||
state: AppState,
|
||||
context: RequestContext,
|
||||
query: StreamSnapshotQuery,
|
||||
payload: Value,
|
||||
) {
|
||||
let _ = socket
|
||||
.send(Message::Text(payload.to_string().into()))
|
||||
.send(serialize_snapshot_message(&payload))
|
||||
.await;
|
||||
|
||||
while let Some(message) = socket.next().await {
|
||||
@@ -34,17 +40,41 @@ async fn handle_socket(mut socket: WebSocket, context: RequestContext) {
|
||||
|
||||
match message {
|
||||
Message::Text(text) => {
|
||||
let echo = json!({
|
||||
"kind": "ws_echo",
|
||||
"traceId": context.trace.trace_id,
|
||||
"text": text.to_string(),
|
||||
});
|
||||
if socket
|
||||
.send(Message::Text(echo.to_string().into()))
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
break;
|
||||
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,
|
||||
@@ -52,3 +82,62 @@ async fn handle_socket(mut socket: WebSocket, context: RequestContext) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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\""));
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user