fix(stub): resolve TCP socket leak on target server disconnect and hanging connection attempt
This commit is contained in:
1 parent
8e10950fc0
commit
1f40e6d32b
3 files changed
+89
-7
No files matched your search
@@ -22,9 +22,8 @@ pub async fn read_mcp_message<R: tokio::io::AsyncRead + Unpin>(
|
|||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
|
|
||||||
let lower_line = line.to_lowercase();
|
if line.len() >= 15 && line[..15].eq_ignore_ascii_case("content-length:") {
|
||||||
if let Some(len_str) = lower_line.strip_prefix("content-length:") {
|
length = line[15..].trim().parse().unwrap_or(0);
|
||||||
length = len_str.trim().parse().unwrap_or(0);
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -39,3 +38,4 @@ pub async fn read_mcp_message<R: tokio::io::AsyncRead + Unpin>(
|
|||||||
|
|
||||||
String::from_utf8(buffer).ok()
|
String::from_utf8(buffer).ok()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -0,0 +1,66 @@
|
|||||||
|
import os
|
||||||
|
|
||||||
|
def fix_stub_leaks():
|
||||||
|
filepath = 'stub/src/main.rs'
|
||||||
|
with open(filepath, 'r', encoding='utf-8') as f:
|
||||||
|
content = f.read()
|
||||||
|
|
||||||
|
# 1. Fix connect_async to handle shutdown and timeout
|
||||||
|
old_connect = """ match tokio_tungstenite::connect_async(request).await {"""
|
||||||
|
new_connect = """ let connect_result = tokio::select! {
|
||||||
|
_ = shutdown_rx.recv() => {
|
||||||
|
tracing::info!("Shutdown received during connect");
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
res = tokio::time::timeout(
|
||||||
|
tokio::time::Duration::from_secs(5),
|
||||||
|
tokio_tungstenite::connect_async(request)
|
||||||
|
) => res,
|
||||||
|
};
|
||||||
|
|
||||||
|
match connect_result {
|
||||||
|
Ok(Ok((ws_stream, _))) => {"""
|
||||||
|
|
||||||
|
if old_connect in content:
|
||||||
|
content = content.replace(old_connect, new_connect)
|
||||||
|
print("Replaced connect_async")
|
||||||
|
|
||||||
|
# Fix Err block to match the new match structure
|
||||||
|
old_err = """ Err(e) => {
|
||||||
|
tracing::error!("Failed to connect via WSS: {}", e);
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
||||||
|
}"""
|
||||||
|
new_err = """ Ok(Err(e)) => {
|
||||||
|
tracing::error!("Failed to connect via WSS: {}", e);
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
||||||
|
}
|
||||||
|
Err(_) => {
|
||||||
|
tracing::error!("Connection attempt timed out");
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
||||||
|
}"""
|
||||||
|
|
||||||
|
if old_err in content:
|
||||||
|
content = content.replace(old_err, new_err)
|
||||||
|
print("Replaced Err branch")
|
||||||
|
|
||||||
|
# 2. Fix the break in send_task that exits the stub instead of reconnecting
|
||||||
|
old_select_send = """ _ = &mut send_task => {
|
||||||
|
tracing::error!("Send task exited");
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
|
||||||
|
recv_task.abort();
|
||||||
|
break;
|
||||||
|
}"""
|
||||||
|
new_select_send = """ _ = &mut send_task => {
|
||||||
|
tracing::error!("Send task exited");
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
|
||||||
|
recv_task.abort();
|
||||||
|
}"""
|
||||||
|
|
||||||
|
if old_select_send in content:
|
||||||
|
content = content.replace(old_select_send, new_select_send)
|
||||||
|
print("Replaced select send_task")
|
||||||
|
|
||||||
|
with open(filepath, 'w', encoding='utf-8') as f:
|
||||||
|
f.write(content)
|
||||||
|
|
||||||
|
fix_stub_leaks()
|
||||||
+20
-4
@@ -74,8 +74,19 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
match tokio_tungstenite::connect_async(request).await {
|
let connect_result = tokio::select! {
|
||||||
Ok((ws_stream, _)) => {
|
_ = shutdown_rx.recv() => {
|
||||||
|
tracing::info!("Shutdown received during connect");
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
res = tokio::time::timeout(
|
||||||
|
tokio::time::Duration::from_secs(5),
|
||||||
|
tokio_tungstenite::connect_async(request)
|
||||||
|
) => res,
|
||||||
|
};
|
||||||
|
|
||||||
|
match connect_result {
|
||||||
|
Ok(Ok((ws_stream, _))) => {
|
||||||
tracing::info!("Successfully connected to target server");
|
tracing::info!("Successfully connected to target server");
|
||||||
let (mut write, mut read) = ws_stream.split();
|
let (mut write, mut read) = ws_stream.split();
|
||||||
|
|
||||||
@@ -138,7 +149,6 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
tracing::error!("Send task exited");
|
tracing::error!("Send task exited");
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
|
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
|
||||||
recv_task.abort();
|
recv_task.abort();
|
||||||
break;
|
|
||||||
}
|
}
|
||||||
_ = &mut recv_task => {
|
_ = &mut recv_task => {
|
||||||
tracing::error!("Recv task exited");
|
tracing::error!("Recv task exited");
|
||||||
@@ -147,12 +157,18 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
}
|
}
|
||||||
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Ok(Err(e)) => {
|
||||||
tracing::error!("Failed to connect via WSS: {}", e);
|
tracing::error!("Failed to connect via WSS: {}", e);
|
||||||
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
||||||
}
|
}
|
||||||
|
Err(_) => {
|
||||||
|
tracing::error!("Connection attempt timed out");
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Ok(())
|
Ok(())
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
Reference in new issue
Block a user