diff --git a/server/src/main.rs b/server/src/main.rs index adc3bfc..6b23036 100644 --- a/server/src/main.rs +++ b/server/src/main.rs @@ -247,6 +247,8 @@ fn run_server(state: Arc) -> Result<(), Box> })) })) .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>, + Query(q): Query, + Json(payload): Json, +) -> 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>, +) -> axum::response::sse::Sse>> { + let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst)); + let (tx, rx) = mpsc::channel::(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>,