feat: restore dual transport (websockets and sse)

This commit is contained in:
Riz Ashraf committed 2026-09-12 20:22:51 +01:00
1 parent 8d34634506
commit 04b9388074
1 file changed
+45
+45
View File
@@ -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>>,