put back cancellation token and add stop button support to cancel for non alpha path

This commit is contained in:
Zane Staggs
2025-10-14 18:04:52 -07:00
parent 5feba293c7
commit db58d09f76
7 changed files with 116 additions and 11 deletions
+88 -10
View File
@@ -85,6 +85,8 @@ pub struct ChatRequest {
session_id: String,
recipe_name: Option<String>,
recipe_version: Option<String>,
#[serde(default)]
background_mode: bool,
}
pub struct SseResponse {
@@ -144,7 +146,12 @@ enum MessageEvent {
Ping,
}
async fn stream_event(event: MessageEvent, tx: &mpsc::Sender<String>) -> bool {
async fn stream_event(
event: MessageEvent,
tx: &mpsc::Sender<String>,
cancel_token: &CancellationToken,
background_mode: bool,
) -> bool {
let json = serde_json::to_string(&event).unwrap_or_else(|e| {
format!(
r#"{{"type":"Error","error":"Failed to serialize event: {}"}}"#,
@@ -152,10 +159,14 @@ async fn stream_event(event: MessageEvent, tx: &mpsc::Sender<String>) -> bool {
)
});
if tx.send(format!("data: {}\n\n", json)).await.is_err() {
tracing::info!("client hung up - continuing agent task");
// Don't cancel the agent task when client disconnects
// Agent should continue running and the session will remain in_use=true
// so that when the client reconnects via the session stream, they can see progress
if background_mode {
tracing::info!(
"client disconnected - continuing agent task in background (ALPHA mode)"
);
} else {
tracing::info!("client disconnected - cancelling agent task (backwards compatibility)");
cancel_token.cancel();
}
return false;
}
true
@@ -210,12 +221,19 @@ pub async fn reply(
let cancel_token = CancellationToken::new();
let messages = Conversation::new_unvalidated(request.messages);
let background_mode = request.background_mode;
if background_mode {
let mut tokens = state.session_cancellation_tokens.lock().await;
tokens.insert(session_id.clone(), cancel_token.clone());
}
let task_cancel = cancel_token.clone();
let task_tx = tx.clone();
let task_state = state.clone();
drop(tokio::spawn(async move {
let agent = match state.get_agent(session_id.clone()).await {
let agent = match task_state.get_agent(session_id.clone()).await {
Ok(agent) => agent,
Err(e) => {
tracing::error!("Failed to get session agent: {}", e);
@@ -224,6 +242,8 @@ pub async fn reply(
error: format!("Failed to get session agent: {}", e),
},
&task_tx,
&task_cancel,
background_mode,
)
.await;
return;
@@ -239,6 +259,8 @@ pub async fn reply(
error: format!("Failed to read session: {}", e),
},
&task_tx,
&task_cancel,
background_mode,
)
.await;
return;
@@ -270,6 +292,8 @@ pub async fn reply(
error: e.to_string(),
},
&task_tx,
&task_cancel,
background_mode,
)
.await;
return;
@@ -288,7 +312,7 @@ pub async fn reply(
}
_ = heartbeat_interval.tick() => {
if client_connected {
client_connected = stream_event(MessageEvent::Ping, &tx).await;
client_connected = stream_event(MessageEvent::Ping, &tx, &task_cancel, background_mode).await;
}
}
response = timeout(Duration::from_millis(500), stream.next()) => {
@@ -302,7 +326,7 @@ pub async fn reply(
// Only send message to client if it's user_visible and client is still connected
if message.is_user_visible() && client_connected {
client_connected = stream_event(MessageEvent::Message { message }, &tx).await;
client_connected = stream_event(MessageEvent::Message { message }, &tx, &task_cancel, background_mode).await;
}
}
Ok(Some(Ok(AgentEvent::HistoryReplaced(new_messages)))) => {
@@ -313,7 +337,7 @@ pub async fn reply(
}
Ok(Some(Ok(AgentEvent::ModelChange { model, mode }))) => {
if client_connected {
client_connected = stream_event(MessageEvent::ModelChange { model, mode }, &tx).await;
client_connected = stream_event(MessageEvent::ModelChange { model, mode }, &tx, &task_cancel, background_mode).await;
}
}
Ok(Some(Ok(AgentEvent::McpNotification((request_id, n))))) => {
@@ -321,7 +345,7 @@ pub async fn reply(
client_connected = stream_event(MessageEvent::Notification{
request_id: request_id.clone(),
message: n,
}, &tx).await;
}, &tx, &task_cancel, background_mode).await;
}
}
@@ -333,6 +357,8 @@ pub async fn reply(
error: e.to_string(),
},
&tx,
&task_cancel,
background_mode,
).await;
}
break;
@@ -406,8 +432,15 @@ pub async fn reply(
reason: "stop".to_string(),
},
&task_tx,
&task_cancel,
background_mode,
)
.await;
if background_mode {
let mut tokens = task_state.session_cancellation_tokens.lock().await;
tokens.remove(&session_id);
}
}));
Ok(SseResponse::new(stream))
}
@@ -459,6 +492,49 @@ pub async fn confirm_permission(
Ok(Json(Value::Object(serde_json::Map::new())))
}
#[derive(Debug, Deserialize, Serialize, ToSchema)]
pub struct CancelRequest {
session_id: String,
}
#[utoipa::path(
post,
path = "/cancel",
request_body = CancelRequest,
responses(
(status = 200, description = "Agent task cancelled successfully"),
(status = 404, description = "Session not found"),
(status = 500, description = "Internal server error")
)
)]
pub async fn cancel_agent(
State(state): State<Arc<AppState>>,
Json(request): Json<CancelRequest>,
) -> Result<Json<Value>, StatusCode> {
tracing::info!("Cancelling agent task for session: {}", request.session_id);
// Get the cancellation token for this session and trigger it
let mut tokens = state.session_cancellation_tokens.lock().await;
if let Some(cancel_token) = tokens.remove(&request.session_id) {
cancel_token.cancel();
tracing::info!("Agent task cancelled for session: {}", request.session_id);
Ok(Json(serde_json::json!({
"status": "cancelled",
"session_id": request.session_id
})))
} else {
tracing::warn!(
"No active cancellation token found for session: {}",
request.session_id
);
// Still return success - the session may have already completed
Ok(Json(serde_json::json!({
"status": "not_active",
"session_id": request.session_id
})))
}
}
pub fn routes(state: Arc<AppState>) -> Router {
Router::new()
.route(
@@ -466,6 +542,7 @@ pub fn routes(state: Arc<AppState>) -> Router {
post(reply).layer(DefaultBodyLimit::max(50 * 1024 * 1024)),
)
.route("/confirm", post(confirm_permission))
.route("/cancel", post(cancel_agent))
.with_state(state)
}
@@ -496,6 +573,7 @@ mod tests {
session_id: "test-session".to_string(),
recipe_name: None,
recipe_version: None,
background_mode: false,
})
.unwrap(),
))
+5
View File
@@ -6,6 +6,8 @@ use std::path::PathBuf;
use std::sync::atomic::AtomicUsize;
use std::sync::Arc;
use tokio::sync::Mutex;
use tokio_util::sync::CancellationToken;
#[derive(Clone)]
pub struct AppState {
pub(crate) agent_manager: Arc<AgentManager>,
@@ -13,6 +15,8 @@ pub struct AppState {
pub session_counter: Arc<AtomicUsize>,
/// Tracks sessions that have already emitted recipe telemetry to prevent double counting.
recipe_session_tracker: Arc<Mutex<HashSet<String>>>,
/// Tracks cancellation tokens for active sessions
pub(crate) session_cancellation_tokens: Arc<Mutex<HashMap<String, CancellationToken>>>,
}
impl AppState {
@@ -23,6 +27,7 @@ impl AppState {
recipe_file_hash_map: Arc::new(Mutex::new(HashMap::new())),
session_counter: Arc::new(AtomicUsize::new(0)),
recipe_session_tracker: Arc::new(Mutex::new(HashSet::new())),
session_cancellation_tokens: Arc::new(Mutex::new(HashMap::new())),
}))
}
+3
View File
@@ -2052,6 +2052,9 @@
"session_id"
],
"properties": {
"background_mode": {
"type": "boolean"
},
"messages": {
"type": "array",
"items": {
+1
View File
@@ -26,6 +26,7 @@ export type AuthorRequest = {
};
export type ChatRequest = {
background_mode?: boolean;
messages: Array<Message>;
recipe_name?: string | null;
recipe_version?: string | null;
+17 -1
View File
@@ -282,6 +282,22 @@ export const useChatEngine = ({
setLastInteractionTime(Date.now());
stopPowerSaveBlocker();
// In ALPHA mode, when user explicitly stops, we need to cancel the agent task
// This is different from a page refresh where we want the agent to continue
if (process.env.ALPHA === 'true' && chat.sessionId) {
fetch(getApiUrl('/cancel'), {
method: 'POST',
headers: {
'Content-Type': 'application/json',
},
body: JSON.stringify({
session_id: chat.sessionId,
}),
}).catch((error) => {
console.error('Failed to cancel agent task:', error);
});
}
// Handle stopping the message stream
const lastMessage = messages[messages.length - 1];
@@ -361,7 +377,7 @@ export const useChatEngine = ({
setMessages([...messages, responseMessage]);
}
}
}, [stop, messages, _setInput, setMessages, stopPowerSaveBlocker]);
}, [stop, messages, _setInput, setMessages, stopPowerSaveBlocker, chat.sessionId]);
// Since server now handles all filtering, we just use messages directly
const filteredMessages = useMemo(() => {
+1
View File
@@ -84,6 +84,7 @@ export function useChatStream({
body: JSON.stringify({
session_id: sessionIdRef.current,
messages: currentMessages,
background_mode: true,
}),
signal: abortControllerRef.current.signal,
});
+1
View File
@@ -408,6 +408,7 @@ export function useMessageStream({
},
body: JSON.stringify({
messages: requestMessages,
background_mode: true,
...extraMetadataRef.current.body,
}),
signal: abortController.signal,