feat: restore dual transport (websockets and sse)
This commit is contained in:
1 parent
8d34634506
commit
04b9388074
1 file changed
+45
@@ -247,6 +247,8 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
|||||||
}))
|
}))
|
||||||
}))
|
}))
|
||||||
.route("/ws", get(ws_handler))
|
.route("/ws", get(ws_handler))
|
||||||
|
.route("/sse", get(sse_handler))
|
||||||
|
.route("/messages", post(message_handler))
|
||||||
.route("/health", get(health_handler))
|
.route("/health", get(health_handler))
|
||||||
.route("/gate/verify", get(gate_verify_handler))
|
.route("/gate/verify", get(gate_verify_handler))
|
||||||
.route("/gate/set", post(gate_set_handler))
|
.route("/gate/set", post(gate_set_handler))
|
||||||
@@ -502,6 +504,49 @@ eprintln!("MCP Memory Server running on http://127.0.0.1:3000/sse");
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(serde::Deserialize)]
|
||||||
|
struct MsgQuery {
|
||||||
|
session_id: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn message_handler(
|
||||||
|
State(state): State<Arc<AppState>>,
|
||||||
|
Query(q): Query<MsgQuery>,
|
||||||
|
Json(payload): Json<serde_json::Value>,
|
||||||
|
) -> impl axum::response::IntoResponse {
|
||||||
|
let session_id = q.session_id;
|
||||||
|
if let Some(response) = state.handler.handle_request(payload).await {
|
||||||
|
let res_str = serde_json::to_string(&response).unwrap();
|
||||||
|
let tx_opt = state.clients.read().unwrap().get(&session_id).cloned();
|
||||||
|
if let Some(tx) = tx_opt {
|
||||||
|
let _ = tx.send(res_str).await;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
(axum::http::StatusCode::ACCEPTED, "Accepted").into_response()
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn sse_handler(
|
||||||
|
State(state): State<Arc<AppState>>,
|
||||||
|
) -> axum::response::sse::Sse<impl tokio_stream::Stream<Item = Result<axum::response::sse::Event, std::convert::Infallible>>> {
|
||||||
|
let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst));
|
||||||
|
let (tx, rx) = mpsc::channel::<String>(100);
|
||||||
|
state.clients.write().unwrap().insert(session_id.clone(), tx.clone());
|
||||||
|
|
||||||
|
let endpoint = format!("/messages?session_id={}", session_id);
|
||||||
|
let _ = tx.send(format!("endpoint|{}", endpoint)).await;
|
||||||
|
|
||||||
|
let rx_stream = tokio_stream::wrappers::ReceiverStream::new(rx);
|
||||||
|
let event_stream = rx_stream.map(|msg| {
|
||||||
|
if let Some(ep) = msg.strip_prefix("endpoint|") {
|
||||||
|
Ok(axum::response::sse::Event::default().event("endpoint").data(ep))
|
||||||
|
} else {
|
||||||
|
Ok(axum::response::sse::Event::default().event("message").data(msg))
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
axum::response::sse::Sse::new(event_stream).keep_alive(axum::response::sse::KeepAlive::new())
|
||||||
|
}
|
||||||
|
|
||||||
async fn ws_handler(
|
async fn ws_handler(
|
||||||
ws: WebSocketUpgrade,
|
ws: WebSocketUpgrade,
|
||||||
State(state): State<Arc<AppState>>,
|
State(state): State<Arc<AppState>>,
|
||||||
|
|||||||
Reference in new issue
Block a user