Detect client disconnects and cancel tool calls (#3782)
This commit is contained in:
@@ -99,19 +99,24 @@ enum MessageEvent {
|
||||
request_id: String,
|
||||
message: ServerNotification,
|
||||
},
|
||||
Ping,
|
||||
}
|
||||
|
||||
async fn stream_event(
|
||||
event: MessageEvent,
|
||||
tx: &mpsc::Sender<String>,
|
||||
) -> Result<(), mpsc::error::SendError<String>> {
|
||||
cancel_token: &CancellationToken,
|
||||
) {
|
||||
let json = serde_json::to_string(&event).unwrap_or_else(|e| {
|
||||
format!(
|
||||
r#"{{"type":"Error","error":"Failed to serialize event: {}"}}"#,
|
||||
e
|
||||
)
|
||||
});
|
||||
tx.send(format!("data: {}\n\n", json)).await
|
||||
if tx.send(format!("data: {}\n\n", json)).await.is_err() {
|
||||
tracing::info!("client hung up");
|
||||
cancel_token.cancel();
|
||||
}
|
||||
}
|
||||
|
||||
async fn reply_handler(
|
||||
@@ -144,6 +149,7 @@ async fn reply_handler(
|
||||
error: "No agent configured".to_string(),
|
||||
},
|
||||
&task_tx,
|
||||
&cancel_token,
|
||||
)
|
||||
.await;
|
||||
return;
|
||||
@@ -173,11 +179,12 @@ async fn reply_handler(
|
||||
Ok(stream) => stream,
|
||||
Err(e) => {
|
||||
tracing::error!("Failed to start reply stream: {:?}", e);
|
||||
let _ = stream_event(
|
||||
stream_event(
|
||||
MessageEvent::Error {
|
||||
error: e.to_string(),
|
||||
},
|
||||
&task_tx,
|
||||
&cancel_token,
|
||||
)
|
||||
.await;
|
||||
return;
|
||||
@@ -194,6 +201,7 @@ async fn reply_handler(
|
||||
error: format!("Failed to get session path: {}", e),
|
||||
},
|
||||
&task_tx,
|
||||
&cancel_token,
|
||||
)
|
||||
.await;
|
||||
return;
|
||||
@@ -201,81 +209,61 @@ async fn reply_handler(
|
||||
};
|
||||
let saved_message_count = all_messages.len();
|
||||
|
||||
let mut heartbeat_interval = tokio::time::interval(Duration::from_millis(500));
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = task_cancel.cancelled() => {
|
||||
tracing::info!("Agent task cancelled");
|
||||
_ = task_cancel.cancelled() => {
|
||||
tracing::info!("Agent task cancelled");
|
||||
break;
|
||||
}
|
||||
_ = heartbeat_interval.tick() => {
|
||||
stream_event(MessageEvent::Ping, &tx, &cancel_token).await;
|
||||
}
|
||||
response = timeout(Duration::from_millis(500), stream.next()) => {
|
||||
match response {
|
||||
Ok(Some(Ok(AgentEvent::Message(message)))) => {
|
||||
push_message(&mut all_messages, message.clone());
|
||||
stream_event(MessageEvent::Message { message }, &tx, &cancel_token).await;
|
||||
}
|
||||
Ok(Some(Ok(AgentEvent::HistoryReplaced(new_messages)))) => {
|
||||
// Replace the message history with the compacted messages
|
||||
all_messages = new_messages;
|
||||
// Note: We don't send this as a stream event since it's an internal operation
|
||||
// The client will see the compaction notification message that was sent before this event
|
||||
}
|
||||
Ok(Some(Ok(AgentEvent::ModelChange { model, mode }))) => {
|
||||
stream_event(MessageEvent::ModelChange { model, mode }, &tx, &cancel_token).await;
|
||||
}
|
||||
Ok(Some(Ok(AgentEvent::McpNotification((request_id, n))))) => {
|
||||
stream_event(MessageEvent::Notification{
|
||||
request_id: request_id.clone(),
|
||||
message: n,
|
||||
}, &tx, &cancel_token).await;
|
||||
}
|
||||
|
||||
Ok(Some(Err(e))) => {
|
||||
tracing::error!("Error processing message: {}", e);
|
||||
stream_event(
|
||||
MessageEvent::Error {
|
||||
error: e.to_string(),
|
||||
},
|
||||
&tx,
|
||||
&cancel_token,
|
||||
).await;
|
||||
break;
|
||||
}
|
||||
Ok(None) => {
|
||||
break;
|
||||
}
|
||||
Err(_) => {
|
||||
if tx.is_closed() {
|
||||
break;
|
||||
}
|
||||
response = timeout(Duration::from_millis(500), stream.next()) => {
|
||||
match response {
|
||||
Ok(Some(Ok(AgentEvent::Message(message)))) => {
|
||||
push_message(&mut all_messages, message.clone());
|
||||
if let Err(e) = stream_event(MessageEvent::Message { message }, &tx).await {
|
||||
tracing::error!("Error sending message through channel: {}", e);
|
||||
let _ = stream_event(
|
||||
MessageEvent::Error {
|
||||
error: e.to_string(),
|
||||
},
|
||||
&tx,
|
||||
).await;
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(Some(Ok(AgentEvent::HistoryReplaced(new_messages)))) => {
|
||||
// Replace the message history with the compacted messages
|
||||
all_messages = new_messages;
|
||||
// Note: We don't send this as a stream event since it's an internal operation
|
||||
// The client will see the compaction notification message that was sent before this event
|
||||
}
|
||||
Ok(Some(Ok(AgentEvent::ModelChange { model, mode }))) => {
|
||||
if let Err(e) = stream_event(MessageEvent::ModelChange { model, mode }, &tx).await {
|
||||
tracing::error!("Error sending model change through channel: {}", e);
|
||||
let _ = stream_event(
|
||||
MessageEvent::Error {
|
||||
error: e.to_string(),
|
||||
},
|
||||
&tx,
|
||||
).await;
|
||||
}
|
||||
}
|
||||
Ok(Some(Ok(AgentEvent::McpNotification((request_id, n))))) => {
|
||||
if let Err(e) = stream_event(MessageEvent::Notification{
|
||||
request_id: request_id.clone(),
|
||||
message: n,
|
||||
}, &tx).await {
|
||||
tracing::error!("Error sending message through channel: {}", e);
|
||||
let _ = stream_event(
|
||||
MessageEvent::Error {
|
||||
error: e.to_string(),
|
||||
},
|
||||
&tx,
|
||||
).await;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Some(Err(e))) => {
|
||||
tracing::error!("Error processing message: {}", e);
|
||||
let _ = stream_event(
|
||||
MessageEvent::Error {
|
||||
error: e.to_string(),
|
||||
},
|
||||
&tx,
|
||||
).await;
|
||||
break;
|
||||
}
|
||||
Ok(None) => {
|
||||
break;
|
||||
}
|
||||
Err(_) => {
|
||||
if tx.is_closed() {
|
||||
break;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if all_messages.len() > saved_message_count {
|
||||
@@ -301,6 +289,7 @@ async fn reply_handler(
|
||||
reason: "stop".to_string(),
|
||||
},
|
||||
&task_tx,
|
||||
&cancel_token,
|
||||
)
|
||||
.await;
|
||||
}));
|
||||
|
||||
Reference in New Issue
Block a user