a2a_protocol_server/handler/lifecycle/
subscribe.rs1use std::collections::HashMap;
9use std::time::Instant;
10
11use a2a_protocol_types::params::TaskIdParams;
12use a2a_protocol_types::task::TaskId;
13
14use crate::error::{ServerError, ServerResult};
15use crate::streaming::InMemoryQueueReader;
16
17use super::super::helpers::build_call_context;
18use super::super::RequestHandler;
19
20impl RequestHandler {
21 pub async fn on_resubscribe(
27 &self,
28 params: TaskIdParams,
29 headers: Option<&HashMap<String, String>>,
30 ) -> ServerResult<InMemoryQueueReader> {
31 let start = Instant::now();
32 trace_info!(method = "SubscribeToTask", task_id = %params.id, "handling resubscribe");
33 self.metrics.on_request("SubscribeToTask");
34
35 let tenant = self
36 .resolve_tenant("SubscribeToTask", headers, params.tenant.as_deref())
37 .await?;
38 let result: ServerResult<_> = crate::store::tenant::TenantContext::scope(tenant, async {
39 let call_ctx = build_call_context("SubscribeToTask", headers);
40 self.interceptors.run_before(&call_ctx).await?;
41 self.ensure_required_extensions(&call_ctx)?;
44
45 self.ensure_streaming_supported()?;
49
50 let task_id = TaskId::new(¶ms.id);
51
52 let task = self
54 .task_store
55 .get(&task_id)
56 .await?
57 .ok_or_else(|| ServerError::TaskNotFound(task_id.clone()))?;
58
59 if task.status.state.is_terminal() {
62 return Err(ServerError::UnsupportedOperation(format!(
63 "task {} is in terminal state '{}' and cannot be subscribed to",
64 task_id, task.status.state
65 )));
66 }
67
68 let snapshot = a2a_protocol_types::events::StreamResponse::Task(task);
71 let reader = self
72 .event_queue_manager
73 .subscribe_with_snapshot(&task_id, snapshot.clone())
74 .await
75 .unwrap_or_else(|| InMemoryQueueReader::snapshot_then_end(snapshot));
81
82 self.interceptors.run_after(&call_ctx).await?;
83 Ok(reader)
84 })
85 .await;
86
87 let elapsed = start.elapsed();
88 match &result {
89 Ok(_) => {
90 self.metrics.on_response("SubscribeToTask");
91 self.metrics.on_latency("SubscribeToTask", elapsed);
92 }
93 Err(e) => {
94 self.metrics.on_error("SubscribeToTask", e.metric_label());
95 self.metrics.on_latency("SubscribeToTask", elapsed);
96 }
97 }
98 result
99 }
100}
101
102#[cfg(test)]
103mod tests {
104 use a2a_protocol_types::params::TaskIdParams;
105
106 use crate::agent_executor;
107 use crate::builder::RequestHandlerBuilder;
108 use crate::error::ServerError;
109
110 struct DummyExecutor;
111 agent_executor!(DummyExecutor, |_ctx, _queue| async { Ok(()) });
112
113 #[tokio::test]
114 async fn resubscribe_task_not_found_returns_error() {
115 let handler = RequestHandlerBuilder::new(DummyExecutor).build().unwrap();
116 let params = TaskIdParams {
117 tenant: None,
118 id: "nonexistent-task".to_owned(),
119 };
120 let result = handler.on_resubscribe(params, None).await;
121 assert!(
122 matches!(result, Err(ServerError::TaskNotFound(_))),
123 "expected TaskNotFound for missing task, got: {result:?}"
124 );
125 }
126
127 #[tokio::test]
128 async fn resubscribe_terminal_task_returns_unsupported_operation() {
129 use a2a_protocol_types::task::{ContextId, Task, TaskId, TaskState, TaskStatus};
131
132 let handler = RequestHandlerBuilder::new(DummyExecutor).build().unwrap();
133 let task = Task {
134 id: TaskId::new("t-resub-1"),
135 context_id: ContextId::new("ctx-1"),
136 status: TaskStatus::new(TaskState::Completed),
137 history: None,
138 artifacts: None,
139 metadata: None,
140 };
141 handler.task_store.save(&task).await.unwrap();
142
143 let params = TaskIdParams {
144 tenant: None,
145 id: "t-resub-1".to_owned(),
146 };
147 let result = handler.on_resubscribe(params, None).await;
148 assert!(
149 matches!(result, Err(ServerError::UnsupportedOperation(ref msg)) if msg.contains("terminal")),
150 "expected UnsupportedOperation for terminal task, got: {result:?}"
151 );
152 }
153
154 #[tokio::test]
155 async fn resubscribe_nonterminal_no_queue_returns_snapshot_then_eof() {
156 use crate::streaming::event_queue::EventQueueReader as _;
160 use a2a_protocol_types::task::{ContextId, Task, TaskId, TaskState, TaskStatus};
161
162 let handler = RequestHandlerBuilder::new(DummyExecutor).build().unwrap();
163 let task = Task {
164 id: TaskId::new("t-resub-nonterminal"),
165 context_id: ContextId::new("ctx-1"),
166 status: TaskStatus::new(TaskState::Working),
167 history: None,
168 artifacts: None,
169 metadata: None,
170 };
171 handler.task_store.save(&task).await.unwrap();
172
173 let params = TaskIdParams {
174 tenant: None,
175 id: "t-resub-nonterminal".to_owned(),
176 };
177 let mut reader = handler
178 .on_resubscribe(params, None)
179 .await
180 .expect("resubscribe to a queueless non-terminal task must serve a snapshot stream");
181
182 let first = reader
184 .read()
185 .await
186 .expect("stream must yield the snapshot")
187 .expect("snapshot must not be an error");
188 match first {
189 a2a_protocol_types::events::StreamResponse::Task(t) => {
190 assert_eq!(t.id.0.as_str(), "t-resub-nonterminal");
191 assert_eq!(t.status.state, TaskState::Working);
192 }
193 other => panic!("expected Task snapshot first, got: {other:?}"),
194 }
195
196 assert!(
198 reader.read().await.is_none(),
199 "stream must end cleanly after the snapshot"
200 );
201 }
202
203 #[tokio::test]
204 async fn resubscribe_success_returns_reader() {
205 use a2a_protocol_types::message::{Message, MessageId, MessageRole, Part};
209 use a2a_protocol_types::params::MessageSendParams;
210 use a2a_protocol_types::task::ContextId;
211
212 use crate::handler::SendMessageResult;
213
214 let handler = RequestHandlerBuilder::new(DummyExecutor).build().unwrap();
215
216 let params = MessageSendParams {
218 message: Message {
219 id: MessageId::new("msg-resub"),
220 role: MessageRole::User,
221 parts: vec![Part::text("hello")],
222 context_id: Some(ContextId::new("ctx-resub")),
223 task_id: None,
224 reference_task_ids: None,
225 extensions: None,
226 metadata: None,
227 },
228 configuration: None,
229 metadata: None,
230 tenant: None,
231 };
232
233 let result = handler.on_send_message(params, true, None).await;
234 assert!(matches!(result, Ok(SendMessageResult::Stream(_))));
235
236 let tasks = handler
238 .task_store
239 .list(&a2a_protocol_types::params::ListTasksParams::default())
240 .await
241 .unwrap();
242 assert!(!tasks.tasks.is_empty(), "should have at least one task");
243
244 let task_id = tasks.tasks[0].id.0.clone();
245
246 let sub_params = TaskIdParams {
248 tenant: None,
249 id: task_id,
250 };
251 let sub_result = handler.on_resubscribe(sub_params, None).await;
252 match &sub_result {
256 Ok(_) | Err(ServerError::Internal(_)) => {} Err(e) => panic!("unexpected error: {e:?}"),
258 }
259 }
260
261 #[tokio::test]
262 async fn resubscribe_with_tenant() {
263 let handler = RequestHandlerBuilder::new(DummyExecutor).build().unwrap();
265 let params = TaskIdParams {
266 tenant: Some("test-tenant".to_string()),
267 id: "nonexistent-task".to_owned(),
268 };
269 let result = handler.on_resubscribe(params, None).await;
270 assert!(result.is_err(), "resubscribe for missing task should fail");
271 }
272
273 #[tokio::test]
274 async fn resubscribe_with_headers() {
275 let handler = RequestHandlerBuilder::new(DummyExecutor).build().unwrap();
277 let params = TaskIdParams {
278 tenant: None,
279 id: "nonexistent-task".to_owned(),
280 };
281 let mut headers = std::collections::HashMap::new();
282 headers.insert("authorization".to_string(), "Bearer tok".to_string());
283 let result = handler.on_resubscribe(params, Some(&headers)).await;
284 assert!(result.is_err());
285 }
286
287 #[tokio::test]
288 async fn resubscribe_error_path_records_error_metrics() {
289 use crate::call_context::CallContext;
291 use crate::interceptor::ServerInterceptor;
292 use std::future::Future;
293 use std::pin::Pin;
294
295 struct FailInterceptor;
296 impl ServerInterceptor for FailInterceptor {
297 fn before<'a>(
298 &'a self,
299 _ctx: &'a CallContext,
300 ) -> Pin<Box<dyn Future<Output = a2a_protocol_types::error::A2aResult<()>> + Send + 'a>>
301 {
302 Box::pin(async {
303 Err(a2a_protocol_types::error::A2aError::internal(
304 "forced failure",
305 ))
306 })
307 }
308 fn after<'a>(
309 &'a self,
310 _ctx: &'a CallContext,
311 ) -> Pin<Box<dyn Future<Output = a2a_protocol_types::error::A2aResult<()>> + Send + 'a>>
312 {
313 Box::pin(async { Ok(()) })
314 }
315 }
316
317 let handler = RequestHandlerBuilder::new(DummyExecutor)
318 .with_interceptor(FailInterceptor)
319 .build()
320 .unwrap();
321
322 let params = TaskIdParams {
323 tenant: None,
324 id: "t-resub-fail".to_owned(),
325 };
326 let result = handler.on_resubscribe(params, None).await;
327 assert!(
328 result.is_err(),
329 "resubscribe should fail when interceptor rejects"
330 );
331 }
332}