71 lines
2.1 KiB
Python
71 lines
2.1 KiB
Python
import sys
|
|
|
|
main_path = r"C:\Users\reazul.ashraf\workspace\rust\mcp-memory\server\src\main.rs"
|
|
with open(main_path, "r", encoding="utf-8") as f:
|
|
content = f.read()
|
|
|
|
# 1. Add WebSocket imports
|
|
import_idx = content.find('use axum::{')
|
|
if import_idx == -1:
|
|
print("Could not find axum imports")
|
|
sys.exit(1)
|
|
|
|
import_code = 'use axum::extract::ws::{WebSocketUpgrade, WebSocket, Message};\n'
|
|
content = content[:import_idx] + import_code + content[import_idx:]
|
|
|
|
# 2. Add /ws route
|
|
route_idx = content.find('.route("/sse", get(sse_handler))')
|
|
if route_idx == -1:
|
|
print("Could not find route definition")
|
|
sys.exit(1)
|
|
|
|
route_code = '.route("/ws", get(ws_handler))\n '
|
|
content = content[:route_idx] + route_code + content[route_idx:]
|
|
|
|
# 3. Add ws_handler implementation
|
|
end_idx = len(content)
|
|
|
|
ws_code = """
|
|
async fn ws_handler(
|
|
ws: WebSocketUpgrade,
|
|
State(state): State<Arc<AppState>>,
|
|
) -> impl IntoResponse {
|
|
ws.on_upgrade(move |socket| handle_socket(socket, state))
|
|
}
|
|
|
|
async fn handle_socket(mut socket: WebSocket, state: Arc<AppState>) {
|
|
let mut rx = state.tx.subscribe();
|
|
|
|
// Create a background task to receive messages from the system and send to the WebSocket
|
|
let mut send_task = tokio::spawn(async move {
|
|
while let Ok(msg) = rx.recv().await {
|
|
if socket.send(Message::Text(msg.into())).await.is_err() {
|
|
break;
|
|
}
|
|
}
|
|
});
|
|
|
|
// Create a task to receive messages from the WebSocket (e.g. task completion from UI)
|
|
// In a real app we'd decode JSON-RPC here
|
|
let mut recv_task = tokio::spawn(async move {
|
|
// Just keeping it alive and listening
|
|
// We can add logic to process JSON-RPC from the frontend here later
|
|
loop {
|
|
tokio::time::sleep(tokio::time::Duration::from_secs(3600)).await;
|
|
}
|
|
});
|
|
|
|
tokio::select! {
|
|
_ = (&mut send_task) => recv_task.abort(),
|
|
_ = (&mut recv_task) => send_task.abort(),
|
|
}
|
|
}
|
|
"""
|
|
|
|
content += ws_code
|
|
|
|
with open(main_path, "w", encoding="utf-8") as f:
|
|
f.write(content)
|
|
|
|
print("Injected WS handler!")
|