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("/sse", get(sse_handler))
|
||||
.route("/messages", post(message_handler))
|
||||
.route("/health", get(health_handler))
|
||||
.route("/gate/verify", get(gate_verify_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(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<Arc<AppState>>,
|
||||
|
||||
Reference in new issue
Block a user