a2a_protocol_server/handler/lifecycle/
cancel_task.rs1use std::collections::HashMap;
9use std::time::Instant;
10
11use a2a_protocol_types::params::CancelTaskParams;
12use a2a_protocol_types::task::{Task, TaskId, TaskState, TaskStatus};
13
14use crate::error::{ServerError, ServerResult};
15use crate::request_context::RequestContext;
16
17use super::super::helpers::build_call_context;
18use super::super::RequestHandler;
19
20impl RequestHandler {
21 #[allow(clippy::too_many_lines)]
27 pub async fn on_cancel_task(
28 &self,
29 params: CancelTaskParams,
30 headers: Option<&HashMap<String, String>>,
31 ) -> ServerResult<Task> {
32 let start = Instant::now();
33 trace_info!(method = "CancelTask", task_id = %params.id, "handling cancel task");
34 self.metrics.on_request("CancelTask");
35
36 let tenant = self
37 .resolve_tenant("CancelTask", headers, params.tenant.as_deref())
38 .await?;
39 let result: ServerResult<_> = crate::store::tenant::TenantContext::scope(tenant, async {
40 let call_ctx = build_call_context("CancelTask", headers);
41 self.interceptors.run_before(&call_ctx).await?;
42 self.ensure_required_extensions(&call_ctx)?;
45
46 let task_id = TaskId::new(¶ms.id);
47 let task = self
48 .task_store
49 .get(&task_id)
50 .await?
51 .ok_or_else(|| ServerError::TaskNotFound(task_id.clone()))?;
52
53 if task.status.state.is_terminal() {
54 return Err(ServerError::TaskNotCancelable(task_id));
55 }
56
57 {
59 let tokens = self.cancellation_tokens.read().await;
60 if let Some(entry) = tokens.get(&task_id) {
61 entry.token.cancel();
62 }
63 }
64
65 let ctx = RequestContext::new(
67 a2a_protocol_types::message::Message {
68 id: a2a_protocol_types::message::MessageId::new(
69 uuid::Uuid::new_v4().to_string(),
70 ),
71 role: a2a_protocol_types::message::MessageRole::User,
72 parts: vec![],
73 task_id: Some(task_id.clone()),
74 context_id: Some(task.context_id.clone()),
75 reference_task_ids: None,
76 extensions: None,
77 metadata: None,
78 },
79 task_id.clone(),
80 task.context_id.0.clone(),
81 );
82
83 let writer = self.event_queue_manager.writer_for_cancel(&task_id).await;
91 self.executor.cancel(&ctx, writer.as_ref()).await?;
92
93 let current = self
99 .task_store
100 .get(&task_id)
101 .await?
102 .ok_or_else(|| ServerError::TaskNotFound(task_id.clone()))?;
103 if current.status.state == a2a_protocol_types::task::TaskState::Canceled {
104 self.interceptors.run_after(&call_ctx).await?;
105 return Ok(current);
106 }
107 if current.status.state.is_terminal() {
108 return Err(ServerError::TaskNotCancelable(task_id));
109 }
110
111 let mut updated = current;
112 updated.status = TaskStatus::with_timestamp(TaskState::Canceled);
113 self.task_store.save(&updated).await?;
114 let final_task = self
116 .task_store
117 .get(&task_id)
118 .await?
119 .ok_or_else(|| ServerError::TaskNotFound(task_id.clone()))?;
120
121 self.interceptors.run_after(&call_ctx).await?;
122 Ok(final_task)
123 })
124 .await;
125
126 let elapsed = start.elapsed();
127 match &result {
128 Ok(_) => {
129 self.metrics.on_response("CancelTask");
130 self.metrics.on_latency("CancelTask", elapsed);
131 }
132 Err(e) => {
133 self.metrics.on_error("CancelTask", e.metric_label());
134 self.metrics.on_latency("CancelTask", elapsed);
135 }
136 }
137 result
138 }
139}
140
141#[cfg(test)]
142mod tests {
143 use a2a_protocol_types::params::CancelTaskParams;
144 use a2a_protocol_types::task::{ContextId, Task, TaskId, TaskState, TaskStatus};
145
146 use crate::agent_executor;
147 use crate::builder::RequestHandlerBuilder;
148 use crate::error::ServerError;
149
150 struct DummyExecutor;
151 agent_executor!(DummyExecutor, |_ctx, _queue| async { Ok(()) });
152
153 struct CancelableExecutor;
154 agent_executor!(CancelableExecutor,
155 execute: |_ctx, _queue| async { Ok(()) },
156 cancel: |_ctx, _queue| async { Ok(()) }
157 );
158
159 fn make_completed_task(id: &str) -> Task {
160 Task {
161 id: TaskId::new(id),
162 context_id: ContextId::new("ctx-1"),
163 status: TaskStatus::new(TaskState::Completed),
164 history: None,
165 artifacts: None,
166 metadata: None,
167 }
168 }
169
170 fn make_submitted_task(id: &str) -> Task {
171 Task {
172 id: TaskId::new(id),
173 context_id: ContextId::new("ctx-1"),
174 status: TaskStatus::new(TaskState::Submitted),
175 history: None,
176 artifacts: None,
177 metadata: None,
178 }
179 }
180
181 #[tokio::test]
182 async fn cancel_task_not_found_returns_error() {
183 let handler = RequestHandlerBuilder::new(DummyExecutor).build().unwrap();
184 let params = CancelTaskParams {
185 tenant: None,
186 id: "nonexistent-task".to_owned(),
187 metadata: None,
188 };
189 let result = handler.on_cancel_task(params, None).await;
190 assert!(
191 matches!(result, Err(ServerError::TaskNotFound(_))),
192 "expected TaskNotFound for missing task, got: {result:?}"
193 );
194 }
195
196 #[tokio::test]
197 async fn cancel_task_terminal_state_returns_not_cancelable() {
198 let handler = RequestHandlerBuilder::new(DummyExecutor).build().unwrap();
199 let task = make_completed_task("t-cancel-terminal");
200 handler.task_store.save(&task).await.unwrap();
201
202 let params = CancelTaskParams {
203 tenant: None,
204 id: "t-cancel-terminal".to_owned(),
205 metadata: None,
206 };
207 let result = handler.on_cancel_task(params, None).await;
208 assert!(
209 matches!(result, Err(ServerError::TaskNotCancelable(_))),
210 "expected TaskNotCancelable for completed task, got: {result:?}"
211 );
212 }
213
214 #[tokio::test]
215 async fn cancel_task_non_terminal_succeeds() {
216 let handler = RequestHandlerBuilder::new(CancelableExecutor)
217 .build()
218 .unwrap();
219 let task = make_submitted_task("t-cancel-active");
220 handler.task_store.save(&task).await.unwrap();
221
222 let params = CancelTaskParams {
223 tenant: None,
224 id: "t-cancel-active".to_owned(),
225 metadata: None,
226 };
227 let result = handler.on_cancel_task(params, None).await;
228 assert!(
229 result.is_ok(),
230 "canceling a non-terminal task should succeed, got: {result:?}"
231 );
232 assert_eq!(
233 result.unwrap().status.state,
234 TaskState::Canceled,
235 "canceled task should have Canceled state"
236 );
237 }
238
239 #[tokio::test]
244 async fn cancel_task_does_not_leak_event_queue() {
245 let handler = RequestHandlerBuilder::new(CancelableExecutor)
246 .build()
247 .unwrap();
248 handler
250 .task_store
251 .save(&make_submitted_task("t-no-leak"))
252 .await
253 .unwrap();
254 assert_eq!(
255 handler.event_queue_manager.active_count().await,
256 0,
257 "precondition: no queue exists"
258 );
259
260 let params = CancelTaskParams {
261 tenant: None,
262 id: "t-no-leak".to_owned(),
263 metadata: None,
264 };
265 let result = handler.on_cancel_task(params, None).await;
266 assert!(result.is_ok(), "cancel should succeed, got {result:?}");
267 assert_eq!(result.unwrap().status.state, TaskState::Canceled);
268
269 assert_eq!(
270 handler.event_queue_manager.active_count().await,
271 0,
272 "cancel must not leave a leaked event queue behind"
273 );
274 }
275
276 #[tokio::test]
277 async fn cancel_task_error_path_records_metrics() {
278 let handler = RequestHandlerBuilder::new(DummyExecutor).build().unwrap();
280 let params = CancelTaskParams {
281 tenant: None,
282 id: "nonexistent-for-metrics".to_owned(),
283 metadata: None,
284 };
285 let result = handler.on_cancel_task(params, None).await;
286 assert!(
287 matches!(result, Err(ServerError::TaskNotFound(_))),
288 "expected TaskNotFound, got: {result:?}"
289 );
290 }
292
293 #[tokio::test]
298 async fn cancel_working_task_with_default_executor_succeeds() {
299 let handler = RequestHandlerBuilder::new(DummyExecutor).build().unwrap();
300 let mut task = make_submitted_task("cancel-default-1");
301 task.status = TaskStatus::new(TaskState::Working);
302 handler.task_store.save(&task).await.unwrap();
303
304 let result = handler
305 .on_cancel_task(
306 CancelTaskParams {
307 tenant: None,
308 id: "cancel-default-1".to_owned(),
309 metadata: None,
310 },
311 None,
312 )
313 .await
314 .expect("cancel of a WORKING task must succeed with the default executor");
315 assert_eq!(result.status.state, TaskState::Canceled);
316 }
317}