1use std::collections::HashMap;
9use std::sync::Arc;
10use std::time::Instant;
11
12use a2a_protocol_types::events::{StreamResponse, TaskStatusUpdateEvent};
13use a2a_protocol_types::params::MessageSendParams;
14use a2a_protocol_types::responses::SendMessageResponse;
15use a2a_protocol_types::task::{ContextId, Task, TaskId, TaskState, TaskStatus};
16
17use crate::error::{ServerError, ServerResult};
18use crate::request_context::RequestContext;
19use crate::streaming::EventQueueWriter;
20
21use super::helpers::{build_call_context, validate_id, validate_metadata_object};
22use super::{CancellationEntry, RequestHandler, SendMessageResult};
23
24pub const MAX_TASK_HISTORY_MESSAGES: usize = 1024;
30
31fn shape_response_history(task: &mut Task, history_length: Option<u32>) {
41 task.history = match (task.history.take(), history_length) {
42 (Some(msgs), Some(n)) if n > 0 => {
43 let n = n as usize;
44 if msgs.len() > n {
45 Some(msgs[msgs.len() - n..].to_vec())
46 } else {
47 Some(msgs)
48 }
49 }
50 _ => None,
51 };
52}
53
54fn json_byte_len(value: &serde_json::Value) -> serde_json::Result<usize> {
56 struct CountWriter(usize);
57 impl std::io::Write for CountWriter {
58 fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
59 self.0 += buf.len();
60 Ok(buf.len())
61 }
62 fn flush(&mut self) -> std::io::Result<()> {
63 Ok(())
64 }
65 }
66 let mut w = CountWriter(0);
67 serde_json::to_writer(&mut w, value)?;
68 Ok(w.0)
69}
70
71fn second_send_blocked(entry: &CancellationEntry) -> bool {
81 !entry.token.is_cancelled()
82}
83
84fn token_aged(elapsed: std::time::Duration, max_token_age: std::time::Duration) -> bool {
87 elapsed >= max_token_age
88}
89
90const fn evict_aged_token(queue_live: bool) -> bool {
94 !queue_live
95}
96
97fn token_still_evictable(
108 entry: &CancellationEntry,
109 now: Instant,
110 max_token_age: std::time::Duration,
111) -> bool {
112 entry.token.is_cancelled() || token_aged(now.duration_since(entry.created_at), max_token_age)
113}
114
115impl RequestHandler {
116 pub async fn on_send_message(
125 &self,
126 params: MessageSendParams,
127 streaming: bool,
128 headers: Option<&HashMap<String, String>>,
129 ) -> ServerResult<SendMessageResult> {
130 let method_name = if streaming {
131 "SendStreamingMessage"
132 } else {
133 "SendMessage"
134 };
135 let start = Instant::now();
136 trace_info!(method = method_name, streaming, "handling send message");
137 self.metrics.on_request(method_name);
138
139 let tenant = self
140 .resolve_tenant(method_name, headers, params.tenant.as_deref())
141 .await?;
142 let result = crate::store::tenant::TenantContext::scope(tenant, async {
143 self.send_message_inner(params, streaming, method_name, headers)
144 .await
145 })
146 .await;
147 let elapsed = start.elapsed();
148 match &result {
149 Ok(_) => {
150 self.metrics.on_response(method_name);
151 self.metrics.on_latency(method_name, elapsed);
152 }
153 Err(e) => {
154 self.metrics.on_error(method_name, e.metric_label());
155 self.metrics.on_latency(method_name, elapsed);
156 }
157 }
158 result
159 }
160
161 #[allow(clippy::too_many_lines)]
164 async fn send_message_inner(
165 &self,
166 params: MessageSendParams,
167 streaming: bool,
168 method_name: &str,
169 headers: Option<&HashMap<String, String>>,
170 ) -> ServerResult<SendMessageResult> {
171 let call_ctx = build_call_context(method_name, headers);
172 self.interceptors.run_before(&call_ctx).await?;
173 self.ensure_required_extensions(&call_ctx)?;
176
177 if streaming {
181 self.ensure_streaming_supported()?;
182 }
183
184 if let Some(ref ctx_id) = params.message.context_id {
186 validate_id(&ctx_id.0, "context_id", self.limits.max_id_length)?;
187 }
188 if let Some(ref task_id) = params.message.task_id {
189 validate_id(&task_id.0, "task_id", self.limits.max_id_length)?;
190 }
191
192 if params.message.parts.is_empty() {
194 return Err(ServerError::InvalidParams(
195 "message must contain at least one part".into(),
196 ));
197 }
198
199 validate_metadata_object(params.message.metadata.as_ref(), "message")?;
205 validate_metadata_object(params.metadata.as_ref(), "request")?;
206 for (i, part) in params.message.parts.iter().enumerate() {
207 validate_metadata_object(part.metadata.as_ref(), &format!("message part {i}"))?;
208 }
209
210 let max_meta = self.limits.max_metadata_size;
213 if let Some(ref meta) = params.message.metadata {
214 let meta_size = json_byte_len(meta).map_err(|_| {
215 ServerError::InvalidParams("message metadata is not serializable".into())
216 })?;
217 if meta_size > max_meta {
218 return Err(ServerError::InvalidParams(format!(
219 "message metadata exceeds maximum size ({meta_size} bytes, max {max_meta})"
220 )));
221 }
222 }
223 if let Some(ref meta) = params.metadata {
224 let meta_size = json_byte_len(meta).map_err(|_| {
225 ServerError::InvalidParams("request metadata is not serializable".into())
226 })?;
227 if meta_size > max_meta {
228 return Err(ServerError::InvalidParams(format!(
229 "request metadata exceeds maximum size ({meta_size} bytes, max {max_meta})"
230 )));
231 }
232 }
233
234 let context_id = if let Some(ref ctx) = params.message.context_id {
240 ctx.0.clone()
241 } else if let Some(ref msg_task_id) = params.message.task_id {
242 match self.task_store.get(msg_task_id).await? {
243 Some(task) => task.context_id.0.clone(),
244 None => return Err(ServerError::TaskNotFound(msg_task_id.clone())),
247 }
248 } else {
249 uuid::Uuid::new_v4().to_string()
250 };
251
252 let context_lock = {
256 let mut locks = self.context_locks.write().await;
257 if locks.len() >= self.limits.max_context_locks {
261 locks.retain(|_, v| Arc::strong_count(v) > 1);
262 }
263 locks.entry(context_id.clone()).or_default().clone()
264 };
265 let context_guard = context_lock.lock().await;
266
267 let stored_task = self.find_task_by_context(&context_id).await?;
269
270 let task_id = if let Some(ref msg_task_id) = params.message.task_id {
274 if let Some(ref stored) = stored_task {
275 if msg_task_id != &stored.id {
276 return Err(ServerError::InvalidParams(
277 "message task_id does not match task found for context".into(),
278 ));
279 }
280 if stored.status.state.is_terminal() {
284 return Err(ServerError::UnsupportedOperation(format!(
285 "task {} is in terminal state '{}' and cannot accept new messages",
286 stored.id, stored.status.state
287 )));
288 }
289 } else {
291 let exists = self.task_store.get(msg_task_id).await?.is_some();
295 if !exists {
296 return Err(ServerError::TaskNotFound(msg_task_id.clone()));
297 }
298 return Err(ServerError::InvalidParams(
300 "task_id exists but belongs to a different context".into(),
301 ));
302 }
303 msg_task_id.clone()
304 } else {
305 TaskId::new(uuid::Uuid::new_v4().to_string())
309 };
310
311 let return_immediately = params
313 .configuration
314 .as_ref()
315 .and_then(|c| c.return_immediately)
316 .unwrap_or(false);
317 let response_history_length = params.configuration.as_ref().and_then(|c| c.history_length);
318
319 let use_background = streaming || return_immediately;
324
325 {
334 let tokens = self.cancellation_tokens.read().await;
335 if let Some(entry) = tokens.get(&task_id) {
336 if second_send_blocked(entry) {
337 return Err(ServerError::UnsupportedOperation(format!(
338 "task {task_id} is already being processed; \
339 wait for it to reach input-required or a terminal state before sending again"
340 )));
341 }
342 }
343 }
344
345 trace_debug!(
347 task_id = %task_id,
348 context_id = %context_id,
349 "creating task"
350 );
351 let mut history = stored_task
358 .as_ref()
359 .and_then(|s| s.history.clone())
360 .unwrap_or_default();
361 history.push(params.message.clone());
362 if history.len() > MAX_TASK_HISTORY_MESSAGES {
363 let excess = history.len() - MAX_TASK_HISTORY_MESSAGES;
364 history.drain(..excess);
365 }
366 let task = Task {
367 id: task_id.clone(),
368 context_id: ContextId::new(&context_id),
369 status: TaskStatus::with_timestamp(TaskState::Submitted),
370 history: Some(history),
371 artifacts: stored_task.as_ref().and_then(|s| s.artifacts.clone()),
372 metadata: stored_task.as_ref().and_then(|s| s.metadata.clone()),
373 };
374
375 let mut ctx = RequestContext::new(params.message, task_id.clone(), context_id);
378 if let Some(stored) = stored_task {
379 ctx = ctx.with_stored_task(stored);
380 }
381 if let Some(meta) = params.metadata {
382 ctx = ctx.with_metadata(meta);
383 }
384
385 let (writer, reader, persistence_rx) = match self
392 .event_queue_manager
393 .lease(&task_id, use_background)
394 .await
395 {
396 crate::streaming::QueueLease::Created {
397 writer,
398 reader,
399 persistence_rx,
400 } => (writer, reader, persistence_rx),
401 crate::streaming::QueueLease::Existing => {
402 return Err(ServerError::UnsupportedOperation(format!(
412 "task {task_id} is already being processed; wait for it to reach \
413 input-required or a terminal state before sending again"
414 )));
415 }
416 crate::streaming::QueueLease::CapacityExhausted => {
417 let cap = self
418 .event_queue_manager
419 .max_concurrent_queues()
420 .map_or_else(String::new, |n| format!(" ({n})"));
421 return Err(ServerError::Overloaded(format!(
422 "server at maximum concurrent stream capacity{cap}; retry later"
423 )));
424 }
425 };
426
427 {
432 let (cancelled_ids, aged_candidates): (Vec<TaskId>, Vec<TaskId>) = {
441 let tokens = self.cancellation_tokens.read().await;
442 if tokens.len() >= self.limits.max_cancellation_tokens {
443 let now = Instant::now();
444 let mut cancelled = Vec::new();
445 let mut aged = Vec::new();
446 for (id, entry) in tokens.iter() {
447 if entry.token.is_cancelled() {
448 cancelled.push(id.clone());
449 } else if token_aged(
450 now.duration_since(entry.created_at),
451 self.limits.max_token_age,
452 ) {
453 aged.push(id.clone());
454 }
455 }
456 drop(tokens);
457 (cancelled, aged)
458 } else {
459 (Vec::new(), Vec::new())
460 }
461 };
462
463 let mut stale_ids = cancelled_ids;
468 for id in aged_candidates {
469 let queue_live = self.event_queue_manager.has_queue(&id).await;
470 if evict_aged_token(queue_live) {
471 stale_ids.push(id);
472 }
473 }
474
475 if !stale_ids.is_empty() {
480 let now = Instant::now();
481 let mut tokens = self.cancellation_tokens.write().await;
482 for id in &stale_ids {
483 let evict = tokens
484 .get(id)
485 .is_some_and(|e| token_still_evictable(e, now, self.limits.max_token_age));
486 if evict {
487 tokens.remove(id);
488 }
489 }
490 }
491
492 let mut tokens = self.cancellation_tokens.write().await;
494 tokens.insert(
495 task_id.clone(),
496 CancellationEntry {
497 token: ctx.cancellation_token.clone(),
498 created_at: Instant::now(),
499 },
500 );
501 }
502
503 if let Err(e) = self.task_store.save(&task).await {
506 self.event_queue_manager.destroy(&task_id).await;
507 self.cancellation_tokens.write().await.remove(&task_id);
508 return Err(e.into());
509 }
510
511 drop(context_guard);
514
515 let executor = Arc::clone(&self.executor);
519 let task_id_for_cleanup = task_id.clone();
520 let event_queue_mgr = self.event_queue_manager.clone();
521 let cancel_tokens = Arc::clone(&self.cancellation_tokens);
522 let executor_timeout = self.executor_timeout;
523 let executor_handle = tokio::spawn(async move {
524 trace_debug!(task_id = %ctx.task_id, "executor started");
525
526 #[allow(clippy::items_after_statements)]
531 struct CleanupGuard {
532 task_id: Option<TaskId>,
533 queue_mgr: crate::streaming::EventQueueManager,
534 tokens: std::sync::Arc<tokio::sync::RwLock<HashMap<TaskId, CancellationEntry>>>,
535 }
536 #[allow(clippy::items_after_statements)]
537 impl Drop for CleanupGuard {
538 fn drop(&mut self) {
539 if let Some(tid) = self.task_id.take() {
540 let qmgr = self.queue_mgr.clone();
541 let tokens = std::sync::Arc::clone(&self.tokens);
542 tokio::task::spawn(async move {
543 qmgr.destroy(&tid).await;
544 tokens.write().await.remove(&tid);
545 });
546 }
547 }
548 }
549 let mut cleanup_guard = CleanupGuard {
550 task_id: Some(task_id_for_cleanup.clone()),
551 queue_mgr: event_queue_mgr.clone(),
552 tokens: Arc::clone(&cancel_tokens),
553 };
554
555 let result = {
557 let exec_future = if let Some(timeout) = executor_timeout {
558 tokio::time::timeout(timeout, executor.execute(&ctx, writer.as_ref()))
559 .await
560 .unwrap_or_else(|_| {
561 Err(a2a_protocol_types::error::A2aError::internal(format!(
562 "executor timed out after {}s",
563 timeout.as_secs()
564 )))
565 })
566 } else {
567 executor.execute(&ctx, writer.as_ref()).await
568 };
569 exec_future
570 };
571
572 if let Err(ref e) = result {
573 trace_error!(task_id = %ctx.task_id, error = %e, "executor failed");
574 let fail_event = StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
576 task_id: ctx.task_id.clone(),
577 context_id: ContextId::new(ctx.context_id.clone()),
578 status: TaskStatus::with_timestamp(TaskState::Failed),
579 metadata: Some(serde_json::json!({ "error": e.to_string() })),
580 });
581 if let Err(_write_err) = writer.write(fail_event).await {
582 trace_error!(
583 task_id = %ctx.task_id,
584 error = %_write_err,
585 "failed to write failure event to queue"
586 );
587 }
588 }
589 drop(writer);
591 event_queue_mgr.destroy(&task_id_for_cleanup).await;
594 cancel_tokens.write().await.remove(&task_id_for_cleanup);
595 cleanup_guard.task_id = None;
596 });
597
598 self.interceptors.run_after(&call_ctx).await?;
599
600 if use_background {
601 self.spawn_background_event_processor(
617 task_id.clone(),
618 executor_handle,
619 persistence_rx,
620 task.clone(),
621 );
622
623 if streaming {
624 let mut reader = reader;
627 let mut snapshot = task.clone();
628 shape_response_history(&mut snapshot, response_history_length);
629 reader.set_first_event(StreamResponse::Task(snapshot));
630 Ok(SendMessageResult::Stream(reader))
631 } else {
632 drop(reader);
636 let mut task = task;
637 shape_response_history(&mut task, response_history_length);
638 Ok(SendMessageResult::Response(SendMessageResponse::Task(task)))
639 }
640 } else {
641 let mut final_task = self
645 .collect_events(reader, task_id.clone(), executor_handle)
646 .await?;
647 shape_response_history(&mut final_task, response_history_length);
648 Ok(SendMessageResult::Response(SendMessageResponse::Task(
649 final_task,
650 )))
651 }
652 }
653}
654
655#[cfg(test)]
656mod tests {
657 use super::*;
658 use a2a_protocol_types::message::{Message, MessageId, MessageRole, Part};
659 use a2a_protocol_types::params::{MessageSendParams, SendMessageConfiguration};
660 use a2a_protocol_types::task::ContextId;
661
662 use crate::agent_executor;
663 use crate::builder::RequestHandlerBuilder;
664
665 struct DummyExecutor;
666 agent_executor!(DummyExecutor, |_ctx, _queue| async { Ok(()) });
667
668 fn make_handler() -> RequestHandler {
669 RequestHandlerBuilder::new(DummyExecutor)
670 .build()
671 .expect("default build should succeed")
672 }
673
674 fn make_params(context_id: Option<&str>) -> MessageSendParams {
675 MessageSendParams {
676 message: Message {
677 id: MessageId::new("msg-1"),
678 role: MessageRole::User,
679 parts: vec![Part::text("hello")],
680 context_id: context_id.map(ContextId::new),
681 task_id: None,
682 reference_task_ids: None,
683 extensions: None,
684 metadata: None,
685 },
686 configuration: None,
687 metadata: None,
688 tenant: None,
689 }
690 }
691
692 #[tokio::test]
693 async fn empty_message_parts_returns_invalid_params() {
694 let handler = make_handler();
695 let mut params = make_params(None);
696 params.message.parts = vec![];
697
698 let result = handler.on_send_message(params, false, None).await;
699
700 assert!(
701 matches!(result, Err(ServerError::InvalidParams(_))),
702 "expected InvalidParams for empty parts"
703 );
704 }
705
706 #[tokio::test]
707 async fn oversized_message_metadata_returns_invalid_params() {
708 let handler = make_handler();
709 let mut params = make_params(None);
710 let big_value = "x".repeat(1_100_000);
712 params.message.metadata = Some(serde_json::json!(big_value));
713
714 let result = handler.on_send_message(params, false, None).await;
715
716 assert!(
717 matches!(result, Err(ServerError::InvalidParams(_))),
718 "expected InvalidParams for oversized message metadata"
719 );
720 }
721
722 #[tokio::test]
723 async fn oversized_request_metadata_returns_invalid_params() {
724 let handler = make_handler();
725 let mut params = make_params(None);
726 let big_value = "x".repeat(1_100_000);
728 params.metadata = Some(serde_json::json!(big_value));
729
730 let result = handler.on_send_message(params, false, None).await;
731
732 assert!(
733 matches!(result, Err(ServerError::InvalidParams(_))),
734 "expected InvalidParams for oversized request metadata"
735 );
736 }
737
738 #[tokio::test]
739 async fn non_object_message_metadata_returns_invalid_params() {
740 let handler = make_handler();
743 let mut params = make_params(None);
744 params.message.metadata = Some(serde_json::json!([1, 2, 3]));
745
746 let result = handler.on_send_message(params, false, None).await;
747 assert!(
748 matches!(result, Err(ServerError::InvalidParams(ref msg))
749 if msg.contains("JSON object") && msg.contains("array")),
750 "expected InvalidParams naming the offending kind (array), got: {result:?}"
751 );
752 }
753
754 #[tokio::test]
755 async fn scalar_request_metadata_returns_invalid_params() {
756 let handler = make_handler();
757 let mut params = make_params(None);
758 params.metadata = Some(serde_json::json!("a bare string"));
759
760 let result = handler.on_send_message(params, false, None).await;
761 assert!(
762 matches!(result, Err(ServerError::InvalidParams(ref msg))
763 if msg.contains("JSON object") && msg.contains("string")),
764 "expected InvalidParams naming the offending kind (string), got: {result:?}"
765 );
766 }
767
768 #[tokio::test]
769 async fn non_object_part_metadata_returns_invalid_params() {
770 let handler = make_handler();
771 let mut params = make_params(None);
772 params.message.parts[0].metadata = Some(serde_json::json!(42));
773
774 let result = handler.on_send_message(params, false, None).await;
775 assert!(
776 matches!(result, Err(ServerError::InvalidParams(ref msg))
777 if msg.contains("part 0") && msg.contains("number")),
778 "expected InvalidParams naming the part index and kind (number), got: {result:?}"
779 );
780 }
781
782 #[tokio::test]
783 async fn object_metadata_is_accepted() {
784 let handler = make_handler();
786 let mut params = make_params(None);
787 params.message.metadata = Some(serde_json::json!({"k": "v"}));
788 params.metadata = Some(serde_json::json!({"trace": 1}));
789
790 let result = handler.on_send_message(params, false, None).await;
791 assert!(
792 result.is_ok(),
793 "object metadata must be accepted, got: {result:?}"
794 );
795 }
796
797 #[tokio::test]
798 async fn valid_message_returns_ok() {
799 let handler = make_handler();
800 let params = make_params(None);
801
802 let result = handler.on_send_message(params, false, None).await;
803
804 let send_result = result.expect("expected Ok for valid message");
805 assert!(
806 matches!(
807 send_result,
808 SendMessageResult::Response(SendMessageResponse::Task(_))
809 ),
810 "expected Response(Task) for non-streaming send"
811 );
812 }
813
814 #[tokio::test]
815 async fn return_immediately_returns_task() {
816 let handler = make_handler();
817 let mut params = make_params(None);
818 params.configuration = Some(SendMessageConfiguration {
819 accepted_output_modes: vec!["text/plain".into()],
820 task_push_notification_config: None,
821 history_length: None,
822 return_immediately: Some(true),
823 });
824
825 let result = handler.on_send_message(params, false, None).await;
826
827 assert!(
828 matches!(
829 result,
830 Ok(SendMessageResult::Response(SendMessageResponse::Task(_)))
831 ),
832 "expected Response(Task) for return_immediately=true"
833 );
834 }
835
836 struct CompletingExecutor;
838 agent_executor!(CompletingExecutor, |ctx, queue| async {
839 for state in [TaskState::Working, TaskState::Completed] {
840 let ev = StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
841 task_id: ctx.task_id.clone(),
842 context_id: ContextId::new(ctx.context_id.clone()),
843 status: TaskStatus::with_timestamp(state),
844 metadata: None,
845 });
846 let _ = queue.write(ev).await;
847 }
848 Ok(())
849 });
850
851 struct BlockingExecutor;
854 agent_executor!(BlockingExecutor, |_ctx, _queue| async {
855 tokio::time::sleep(std::time::Duration::from_secs(30)).await;
856 Ok(())
857 });
858
859 async fn poll_task_state(
860 handler: &RequestHandler,
861 task_id: &TaskId,
862 want: TaskState,
863 ) -> TaskState {
864 for _ in 0..200 {
865 if let Ok(Some(t)) = handler.task_store.get(task_id).await {
866 if t.status.state == want {
867 return want;
868 }
869 }
870 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
871 }
872 handler
873 .task_store
874 .get(task_id)
875 .await
876 .ok()
877 .flatten()
878 .map_or(TaskState::Submitted, |t| t.status.state)
879 }
880
881 #[tokio::test]
886 async fn return_immediately_persists_final_state() {
887 let handler = RequestHandlerBuilder::new(CompletingExecutor)
888 .build()
889 .unwrap();
890 let mut params = make_params(Some("ctx-ri"));
891 params.configuration = Some(SendMessageConfiguration {
892 accepted_output_modes: vec!["text/plain".into()],
893 task_push_notification_config: None,
894 history_length: None,
895 return_immediately: Some(true),
896 });
897
898 let SendMessageResult::Response(SendMessageResponse::Task(task)) =
899 handler.on_send_message(params, false, None).await.unwrap()
900 else {
901 panic!("expected an immediate Task response");
902 };
903 assert_eq!(
904 task.status.state,
905 TaskState::Submitted,
906 "snapshot is Submitted"
907 );
908
909 let final_state = poll_task_state(&handler, &task.id, TaskState::Completed).await;
910 assert_eq!(
911 final_state,
912 TaskState::Completed,
913 "fire-and-forget task must reach Completed in the store"
914 );
915 }
916
917 #[tokio::test]
921 async fn concurrent_send_to_in_flight_task_is_rejected() {
922 let handler = RequestHandlerBuilder::new(BlockingExecutor)
923 .build()
924 .unwrap();
925
926 let mut first = make_params(Some("ctx-dup"));
928 first.configuration = Some(SendMessageConfiguration {
929 accepted_output_modes: vec!["text/plain".into()],
930 task_push_notification_config: None,
931 history_length: None,
932 return_immediately: Some(true),
933 });
934 let SendMessageResult::Response(SendMessageResponse::Task(task)) =
935 handler.on_send_message(first, false, None).await.unwrap()
936 else {
937 panic!("expected an immediate Task response");
938 };
939
940 let mut second = make_params(Some("ctx-dup"));
942 second.message.task_id = Some(task.id.clone());
943 let result = handler.on_send_message(second, false, None).await;
944 assert!(
945 matches!(result, Err(ServerError::UnsupportedOperation(_))),
946 "expected rejection of a send to an in-flight task, got {result:?}"
947 );
948 }
949
950 #[tokio::test]
955 async fn stream_cap_exhaustion_returns_overloaded_without_orphan() {
956 let handler = RequestHandlerBuilder::new(BlockingExecutor)
957 .with_max_concurrent_streams(1)
958 .build()
959 .unwrap();
960
961 let mut first = make_params(Some("ctx-a"));
964 first.configuration = Some(SendMessageConfiguration {
965 accepted_output_modes: vec!["text/plain".into()],
966 task_push_notification_config: None,
967 history_length: None,
968 return_immediately: Some(true),
969 });
970 handler.on_send_message(first, false, None).await.unwrap();
971 assert_eq!(handler.event_queue_manager.active_count().await, 1);
972
973 let mut second = make_params(Some("ctx-b"));
975 second.configuration = Some(SendMessageConfiguration {
976 accepted_output_modes: vec!["text/plain".into()],
977 task_push_notification_config: None,
978 history_length: None,
979 return_immediately: Some(true),
980 });
981 let result = handler.on_send_message(second, false, None).await;
982 assert!(
983 matches!(result, Err(ServerError::Overloaded(_))),
984 "expected Overloaded at capacity, got {result:?}"
985 );
986 assert_eq!(
988 handler.event_queue_manager.active_count().await,
989 1,
990 "capacity rejection must not create a queue"
991 );
992 }
993
994 #[tokio::test]
995 async fn empty_context_id_returns_invalid_params() {
996 let handler = make_handler();
997 let params = make_params(Some(""));
998
999 let result = handler.on_send_message(params, false, None).await;
1000
1001 assert!(
1002 matches!(result, Err(ServerError::InvalidParams(_))),
1003 "expected InvalidParams for empty context_id"
1004 );
1005 }
1006
1007 #[tokio::test]
1008 async fn too_long_context_id_returns_invalid_params() {
1009 use crate::handler::limits::HandlerLimits;
1011
1012 let handler = RequestHandlerBuilder::new(DummyExecutor)
1013 .with_handler_limits(HandlerLimits::default().with_max_id_length(10))
1014 .build()
1015 .unwrap();
1016 let long_ctx = "x".repeat(20);
1017 let params = make_params(Some(&long_ctx));
1018
1019 let result = handler.on_send_message(params, false, None).await;
1020 assert!(
1021 matches!(result, Err(ServerError::InvalidParams(ref msg)) if msg.contains("maximum length")),
1022 "expected InvalidParams for too-long context_id"
1023 );
1024 }
1025
1026 #[tokio::test]
1027 async fn too_long_task_id_returns_invalid_params() {
1028 use crate::handler::limits::HandlerLimits;
1030 use a2a_protocol_types::task::TaskId;
1031
1032 let handler = RequestHandlerBuilder::new(DummyExecutor)
1033 .with_handler_limits(HandlerLimits::default().with_max_id_length(10))
1034 .build()
1035 .unwrap();
1036 let mut params = make_params(None);
1037 params.message.task_id = Some(TaskId::new("a".repeat(20)));
1038
1039 let result = handler.on_send_message(params, false, None).await;
1040 assert!(
1041 matches!(result, Err(ServerError::InvalidParams(ref msg)) if msg.contains("maximum length")),
1042 "expected InvalidParams for too-long task_id"
1043 );
1044 }
1045
1046 #[tokio::test]
1047 async fn empty_task_id_returns_invalid_params() {
1048 use a2a_protocol_types::task::TaskId;
1050
1051 let handler = make_handler();
1052 let mut params = make_params(None);
1053 params.message.task_id = Some(TaskId::new(""));
1054
1055 let result = handler.on_send_message(params, false, None).await;
1056 assert!(
1057 matches!(result, Err(ServerError::InvalidParams(ref msg)) if msg.contains("empty")),
1058 "expected InvalidParams for empty task_id"
1059 );
1060 }
1061
1062 #[tokio::test]
1063 async fn task_id_mismatch_returns_invalid_params() {
1064 use a2a_protocol_types::task::{Task, TaskId, TaskState, TaskStatus};
1066
1067 let handler = make_handler();
1068
1069 let task = Task {
1071 id: TaskId::new("stored-task-id"),
1072 context_id: ContextId::new("ctx-existing"),
1073 status: TaskStatus::new(TaskState::InputRequired),
1074 history: None,
1075 artifacts: None,
1076 metadata: None,
1077 };
1078 handler.task_store.save(&task).await.unwrap();
1079
1080 let mut params = make_params(Some("ctx-existing"));
1082 params.message.task_id = Some(TaskId::new("different-task-id"));
1083
1084 let result = handler.on_send_message(params, false, None).await;
1085 assert!(
1086 matches!(result, Err(ServerError::InvalidParams(ref msg)) if msg.contains("does not match")),
1087 "expected InvalidParams for task_id mismatch, got: {result:?}"
1088 );
1089 }
1090
1091 #[tokio::test]
1092 async fn send_message_records_user_message_in_history() {
1093 let handler = make_handler();
1096 let result = handler
1097 .on_send_message(make_params(None), false, None)
1098 .await
1099 .expect("send should succeed");
1100 let task_id = match result {
1101 SendMessageResult::Response(SendMessageResponse::Task(t)) => t.id,
1102 other => panic!("expected task response, got {other:?}"),
1103 };
1104 let stored = handler
1105 .task_store
1106 .get(&task_id)
1107 .await
1108 .expect("get")
1109 .expect("task stored");
1110 let history = stored.history.expect("history populated on send");
1111 assert_eq!(history.len(), 1, "exactly the incoming user message");
1112 assert_eq!(history[0].role, MessageRole::User);
1113 assert_eq!(
1114 history[0].parts[0].text_content(),
1115 Some("hello"),
1116 "history records the message content"
1117 );
1118 }
1119
1120 #[tokio::test]
1121 async fn continuation_appends_history_and_preserves_artifacts() {
1122 use a2a_protocol_types::artifact::Artifact;
1125 let handler = make_handler();
1126 let prior = Task {
1127 id: TaskId::new("cont-task"),
1128 context_id: ContextId::new("ctx-cont"),
1129 status: TaskStatus::new(TaskState::InputRequired),
1130 history: Some(vec![Message {
1131 id: MessageId::new("m-prior"),
1132 role: MessageRole::User,
1133 parts: vec![Part::text("first turn")],
1134 context_id: None,
1135 task_id: None,
1136 reference_task_ids: None,
1137 extensions: None,
1138 metadata: None,
1139 }]),
1140 artifacts: Some(vec![Artifact::new("a1", vec![Part::text("turn-1 output")])]),
1141 metadata: Some(serde_json::json!({"k": "v"})),
1142 };
1143 handler.task_store.save(&prior).await.unwrap();
1144
1145 let mut params = make_params(Some("ctx-cont"));
1146 params.message.task_id = Some(TaskId::new("cont-task"));
1147 handler
1148 .on_send_message(params, false, None)
1149 .await
1150 .expect("continuation should succeed");
1151
1152 let stored = handler
1153 .task_store
1154 .get(&TaskId::new("cont-task"))
1155 .await
1156 .expect("get")
1157 .expect("task stored");
1158 let history = stored.history.expect("history preserved");
1159 assert_eq!(history.len(), 2, "prior message + continuation message");
1160 assert_eq!(history[0].parts[0].text_content(), Some("first turn"));
1161 assert_eq!(history[1].parts[0].text_content(), Some("hello"));
1162 assert!(
1163 stored.artifacts.as_ref().is_some_and(|a| a.len() == 1),
1164 "continuation must not wipe accumulated artifacts"
1165 );
1166 assert_eq!(
1167 stored.metadata,
1168 Some(serde_json::json!({"k": "v"})),
1169 "continuation must not wipe task metadata"
1170 );
1171 }
1172
1173 #[tokio::test]
1174 async fn history_is_capped_at_max_messages() {
1175 let handler = make_handler();
1177 let mut long_history: Vec<Message> = (0..MAX_TASK_HISTORY_MESSAGES)
1178 .map(|i| Message {
1179 id: MessageId::new(format!("m-{i}")),
1180 role: MessageRole::User,
1181 parts: vec![Part::text(format!("msg {i}"))],
1182 context_id: None,
1183 task_id: None,
1184 reference_task_ids: None,
1185 extensions: None,
1186 metadata: None,
1187 })
1188 .collect();
1189 long_history[0].parts = vec![Part::text("OLDEST")];
1190 let prior = Task {
1191 id: TaskId::new("cap-task"),
1192 context_id: ContextId::new("ctx-cap"),
1193 status: TaskStatus::new(TaskState::InputRequired),
1194 history: Some(long_history),
1195 artifacts: None,
1196 metadata: None,
1197 };
1198 handler.task_store.save(&prior).await.unwrap();
1199
1200 let mut params = make_params(Some("ctx-cap"));
1201 params.message.task_id = Some(TaskId::new("cap-task"));
1202 handler
1203 .on_send_message(params, false, None)
1204 .await
1205 .expect("continuation should succeed");
1206
1207 let stored = handler
1208 .task_store
1209 .get(&TaskId::new("cap-task"))
1210 .await
1211 .unwrap()
1212 .unwrap();
1213 let history = stored.history.unwrap();
1214 assert_eq!(history.len(), MAX_TASK_HISTORY_MESSAGES, "capped");
1215 assert_ne!(
1216 history[0].parts[0].text_content(),
1217 Some("OLDEST"),
1218 "the oldest message is dropped first"
1219 );
1220 assert_eq!(
1221 history[MAX_TASK_HISTORY_MESSAGES - 1].parts[0].text_content(),
1222 Some("hello"),
1223 "the newest message is retained"
1224 );
1225 }
1226
1227 #[tokio::test]
1228 async fn send_response_omits_history_by_default_and_honors_history_length() {
1229 use a2a_protocol_types::params::SendMessageConfiguration;
1234 let handler = make_handler();
1235
1236 let result = handler
1237 .on_send_message(make_params(Some("ctx-resp")), false, None)
1238 .await
1239 .expect("send should succeed");
1240 let task = match result {
1241 SendMessageResult::Response(SendMessageResponse::Task(t)) => t,
1242 other => panic!("expected task response, got {other:?}"),
1243 };
1244 assert!(
1245 task.history.is_none(),
1246 "default send response must not echo history"
1247 );
1248 let stored = handler
1249 .task_store
1250 .get(&task.id)
1251 .await
1252 .unwrap()
1253 .expect("task stored");
1254 assert_eq!(
1255 stored.history.as_ref().map(Vec::len),
1256 Some(1),
1257 "the store still keeps the full history"
1258 );
1259
1260 let mut params = make_params(Some("ctx-resp"));
1261 params.message.task_id = Some(task.id.clone());
1262 params.configuration = Some(SendMessageConfiguration {
1263 history_length: Some(10),
1264 ..Default::default()
1265 });
1266 let result = handler
1267 .on_send_message(params, false, None)
1268 .await
1269 .expect("continuation should succeed");
1270 let task = match result {
1271 SendMessageResult::Response(SendMessageResponse::Task(t)) => t,
1272 other => panic!("expected task response, got {other:?}"),
1273 };
1274 assert_eq!(
1275 task.history.as_ref().map(Vec::len),
1276 Some(2),
1277 "historyLength=10 returns the (2) stored messages"
1278 );
1279 }
1280
1281 #[tokio::test]
1282 async fn send_message_with_request_metadata() {
1283 let handler = make_handler();
1285 let mut params = make_params(None);
1286 params.metadata = Some(serde_json::json!({"key": "value"}));
1287
1288 let result = handler.on_send_message(params, false, None).await;
1289 assert!(
1290 result.is_ok(),
1291 "send_message with request metadata should succeed"
1292 );
1293 }
1294
1295 #[tokio::test]
1296 async fn send_message_error_path_records_metrics() {
1297 use crate::call_context::CallContext;
1299 use crate::interceptor::ServerInterceptor;
1300 use std::future::Future;
1301 use std::pin::Pin;
1302
1303 struct FailInterceptor;
1304 impl ServerInterceptor for FailInterceptor {
1305 fn before<'a>(
1306 &'a self,
1307 _ctx: &'a CallContext,
1308 ) -> Pin<Box<dyn Future<Output = a2a_protocol_types::error::A2aResult<()>> + Send + 'a>>
1309 {
1310 Box::pin(async {
1311 Err(a2a_protocol_types::error::A2aError::internal(
1312 "forced failure",
1313 ))
1314 })
1315 }
1316 fn after<'a>(
1317 &'a self,
1318 _ctx: &'a CallContext,
1319 ) -> Pin<Box<dyn Future<Output = a2a_protocol_types::error::A2aResult<()>> + Send + 'a>>
1320 {
1321 Box::pin(async { Ok(()) })
1322 }
1323 }
1324
1325 let handler = RequestHandlerBuilder::new(DummyExecutor)
1326 .with_interceptor(FailInterceptor)
1327 .build()
1328 .unwrap();
1329
1330 let params = make_params(None);
1331 let result = handler.on_send_message(params, false, None).await;
1332 assert!(
1333 result.is_err(),
1334 "send_message should fail when interceptor rejects, exercising error metrics path"
1335 );
1336 }
1337
1338 #[tokio::test]
1339 async fn send_streaming_message_error_path_records_metrics() {
1340 use crate::call_context::CallContext;
1342 use crate::interceptor::ServerInterceptor;
1343 use std::future::Future;
1344 use std::pin::Pin;
1345
1346 struct FailInterceptor;
1347 impl ServerInterceptor for FailInterceptor {
1348 fn before<'a>(
1349 &'a self,
1350 _ctx: &'a CallContext,
1351 ) -> Pin<Box<dyn Future<Output = a2a_protocol_types::error::A2aResult<()>> + Send + 'a>>
1352 {
1353 Box::pin(async {
1354 Err(a2a_protocol_types::error::A2aError::internal(
1355 "forced failure",
1356 ))
1357 })
1358 }
1359 fn after<'a>(
1360 &'a self,
1361 _ctx: &'a CallContext,
1362 ) -> Pin<Box<dyn Future<Output = a2a_protocol_types::error::A2aResult<()>> + Send + 'a>>
1363 {
1364 Box::pin(async { Ok(()) })
1365 }
1366 }
1367
1368 let handler = RequestHandlerBuilder::new(DummyExecutor)
1369 .with_interceptor(FailInterceptor)
1370 .build()
1371 .unwrap();
1372
1373 let params = make_params(None);
1374 let result = handler.on_send_message(params, true, None).await;
1375 assert!(
1376 result.is_err(),
1377 "streaming send_message should fail when interceptor rejects"
1378 );
1379 }
1380
1381 #[tokio::test]
1382 async fn streaming_mode_returns_stream_result() {
1383 let handler = make_handler();
1385 let params = make_params(None);
1386
1387 let result = handler.on_send_message(params, true, None).await;
1388 assert!(
1389 matches!(result, Ok(SendMessageResult::Stream(_))),
1390 "expected Stream result in streaming mode"
1391 );
1392 }
1393
1394 #[tokio::test]
1395 async fn send_message_with_stored_task_continuation() {
1396 use a2a_protocol_types::task::{Task, TaskState, TaskStatus};
1399
1400 let handler = make_handler();
1401
1402 let task = Task {
1404 id: TaskId::new("existing-task"),
1405 context_id: ContextId::new("continue-ctx"),
1406 status: TaskStatus::new(TaskState::InputRequired),
1407 history: None,
1408 artifacts: None,
1409 metadata: None,
1410 };
1411 handler.task_store.save(&task).await.unwrap();
1412
1413 let params = make_params(Some("continue-ctx"));
1415 let result = handler.on_send_message(params, false, None).await;
1416 assert!(
1417 result.is_ok(),
1418 "send_message with existing non-terminal context should succeed"
1419 );
1420 }
1421
1422 #[tokio::test]
1423 async fn send_message_to_terminal_task_returns_unsupported_operation() {
1424 use a2a_protocol_types::task::{Task, TaskState, TaskStatus};
1427
1428 let handler = make_handler();
1429
1430 let task = Task {
1432 id: TaskId::new("done-task"),
1433 context_id: ContextId::new("done-ctx"),
1434 status: TaskStatus::new(TaskState::Completed),
1435 history: None,
1436 artifacts: None,
1437 metadata: None,
1438 };
1439 handler.task_store.save(&task).await.unwrap();
1440
1441 let mut params = make_params(Some("done-ctx"));
1443 params.message.task_id = Some(TaskId::new("done-task"));
1444 let result = handler.on_send_message(params, false, None).await;
1445 assert!(
1446 matches!(result, Err(ServerError::UnsupportedOperation(ref msg)) if msg.contains("terminal")),
1447 "expected UnsupportedOperation for terminal task, got: {result:?}"
1448 );
1449 }
1450
1451 #[tokio::test]
1452 async fn send_message_to_terminal_context_without_task_id_creates_new_task() {
1453 use a2a_protocol_types::task::{Task, TaskState, TaskStatus};
1456
1457 let handler = make_handler();
1458
1459 let task = Task {
1461 id: TaskId::new("old-task"),
1462 context_id: ContextId::new("reuse-ctx"),
1463 status: TaskStatus::new(TaskState::Completed),
1464 history: None,
1465 artifacts: None,
1466 metadata: None,
1467 };
1468 handler.task_store.save(&task).await.unwrap();
1469
1470 let params = make_params(Some("reuse-ctx"));
1472 let result = handler.on_send_message(params, false, None).await;
1473 assert!(
1474 result.is_ok(),
1475 "should create new task on terminal context, got: {result:?}"
1476 );
1477 }
1478
1479 #[tokio::test]
1480 async fn send_message_with_headers() {
1481 let handler = make_handler();
1483 let params = make_params(None);
1484 let mut headers = HashMap::new();
1485 headers.insert("authorization".to_string(), "Bearer test-token".to_string());
1486
1487 let result = handler.on_send_message(params, false, Some(&headers)).await;
1488 let send_result = result.expect("send_message with headers should succeed");
1489 assert!(
1490 matches!(
1491 send_result,
1492 SendMessageResult::Response(SendMessageResponse::Task(_))
1493 ),
1494 "expected Response(Task) for send with headers"
1495 );
1496 }
1497
1498 #[tokio::test]
1499 async fn duplicate_task_id_without_context_match_returns_error() {
1500 use a2a_protocol_types::task::{Task, TaskId as TId, TaskState, TaskStatus};
1502
1503 let handler = make_handler();
1504
1505 let task = Task {
1507 id: TId::new("dup-task"),
1508 context_id: ContextId::new("other-ctx"),
1509 status: TaskStatus::new(TaskState::Completed),
1510 history: None,
1511 artifacts: None,
1512 metadata: None,
1513 };
1514 handler.task_store.save(&task).await.unwrap();
1515
1516 let mut params = make_params(Some("brand-new-ctx"));
1518 params.message.task_id = Some(TId::new("dup-task"));
1519
1520 let result = handler.on_send_message(params, false, None).await;
1521 assert!(
1522 matches!(result, Err(ServerError::InvalidParams(ref msg)) if msg.contains("different context")),
1523 "expected InvalidParams for task_id in different context, got: {result:?}"
1524 );
1525 }
1526
1527 #[tokio::test]
1528 async fn unknown_task_id_returns_task_not_found() {
1529 use a2a_protocol_types::task::TaskId as TId;
1531
1532 let handler = make_handler();
1533
1534 let mut params = make_params(Some("fresh-ctx"));
1536 params.message.task_id = Some(TId::new("nonexistent-task"));
1537
1538 let result = handler.on_send_message(params, false, None).await;
1539 assert!(
1540 matches!(result, Err(ServerError::TaskNotFound(_))),
1541 "expected TaskNotFound for unknown task_id, got: {result:?}"
1542 );
1543 }
1544
1545 #[tokio::test]
1546 async fn send_message_with_tenant() {
1547 let handler = make_handler();
1549 let mut params = make_params(None);
1550 params.tenant = Some("test-tenant".to_string());
1551
1552 let result = handler.on_send_message(params, false, None).await;
1553 let send_result = result.expect("send_message with tenant should succeed");
1554 assert!(
1555 matches!(
1556 send_result,
1557 SendMessageResult::Response(SendMessageResponse::Task(_))
1558 ),
1559 "expected Response(Task) for send with tenant"
1560 );
1561 }
1562
1563 #[tokio::test]
1564 async fn executor_timeout_returns_failed_task() {
1565 use a2a_protocol_types::error::A2aResult;
1567 use std::time::Duration;
1568
1569 struct SlowExecutor;
1570 impl crate::executor::AgentExecutor for SlowExecutor {
1571 fn execute<'a>(
1572 &'a self,
1573 _ctx: &'a crate::request_context::RequestContext,
1574 _queue: &'a dyn crate::streaming::EventQueueWriter,
1575 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = A2aResult<()>> + Send + 'a>>
1576 {
1577 Box::pin(async {
1578 tokio::time::sleep(Duration::from_secs(60)).await;
1579 Ok(())
1580 })
1581 }
1582 }
1583
1584 let handler = RequestHandlerBuilder::new(SlowExecutor)
1585 .with_executor_timeout(Duration::from_millis(50))
1586 .build()
1587 .unwrap();
1588
1589 let params = make_params(None);
1590 let result = handler.on_send_message(params, false, None).await;
1592 assert!(
1594 result.is_ok(),
1595 "executor timeout should still return a task result"
1596 );
1597 }
1598
1599 #[tokio::test]
1600 async fn executor_failure_writes_failed_event() {
1601 use a2a_protocol_types::error::{A2aError, A2aResult};
1603
1604 struct FailExecutor;
1605 impl crate::executor::AgentExecutor for FailExecutor {
1606 fn execute<'a>(
1607 &'a self,
1608 _ctx: &'a crate::request_context::RequestContext,
1609 _queue: &'a dyn crate::streaming::EventQueueWriter,
1610 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = A2aResult<()>> + Send + 'a>>
1611 {
1612 Box::pin(async { Err(A2aError::internal("executor exploded")) })
1613 }
1614 }
1615
1616 let handler = RequestHandlerBuilder::new(FailExecutor).build().unwrap();
1617 let params = make_params(None);
1618
1619 let result = handler.on_send_message(params, false, None).await;
1620 assert!(
1622 result.is_ok(),
1623 "executor failure should produce a task result"
1624 );
1625 }
1626
1627 #[tokio::test]
1628 async fn cancellation_token_sweep_runs_when_map_is_full() {
1629 use crate::handler::limits::HandlerLimits;
1632
1633 struct SlowExec;
1635 impl crate::executor::AgentExecutor for SlowExec {
1636 fn execute<'a>(
1637 &'a self,
1638 _ctx: &'a crate::request_context::RequestContext,
1639 _queue: &'a dyn crate::streaming::EventQueueWriter,
1640 ) -> std::pin::Pin<
1641 Box<
1642 dyn std::future::Future<Output = a2a_protocol_types::error::A2aResult<()>>
1643 + Send
1644 + 'a,
1645 >,
1646 > {
1647 Box::pin(async {
1648 tokio::time::sleep(std::time::Duration::from_secs(10)).await;
1650 Ok(())
1651 })
1652 }
1653 }
1654
1655 let handler = RequestHandlerBuilder::new(SlowExec)
1656 .with_handler_limits(HandlerLimits::default().with_max_cancellation_tokens(2))
1657 .build()
1658 .unwrap();
1659
1660 for _ in 0..3 {
1663 let params = make_params(None);
1664 let _ = handler.on_send_message(params, true, None).await;
1665 }
1666 handler.shutdown().await;
1669 }
1670
1671 #[tokio::test]
1672 async fn stale_cancellation_tokens_cleaned_up() {
1673 use crate::handler::limits::HandlerLimits;
1675 use std::time::Duration;
1676
1677 struct SlowExec2;
1679 impl crate::executor::AgentExecutor for SlowExec2 {
1680 fn execute<'a>(
1681 &'a self,
1682 _ctx: &'a crate::request_context::RequestContext,
1683 _queue: &'a dyn crate::streaming::EventQueueWriter,
1684 ) -> std::pin::Pin<
1685 Box<
1686 dyn std::future::Future<Output = a2a_protocol_types::error::A2aResult<()>>
1687 + Send
1688 + 'a,
1689 >,
1690 > {
1691 Box::pin(async {
1692 tokio::time::sleep(Duration::from_secs(10)).await;
1693 Ok(())
1694 })
1695 }
1696 }
1697
1698 let handler = RequestHandlerBuilder::new(SlowExec2)
1699 .with_handler_limits(
1700 HandlerLimits::default()
1701 .with_max_cancellation_tokens(2)
1702 .with_max_token_age(Duration::from_millis(1)),
1704 )
1705 .build()
1706 .unwrap();
1707
1708 for _ in 0..2 {
1710 let params = make_params(None);
1711 let _ = handler.on_send_message(params, true, None).await;
1712 }
1713
1714 tokio::time::sleep(Duration::from_millis(50)).await;
1716
1717 let params = make_params(None);
1721 let _ = handler.on_send_message(params, true, None).await;
1722
1723 handler.shutdown().await;
1725 }
1726
1727 #[tokio::test]
1728 async fn streaming_executor_failure_writes_error_event() {
1729 use a2a_protocol_types::error::{A2aError, A2aResult};
1731
1732 struct FailExecutor;
1733 impl crate::executor::AgentExecutor for FailExecutor {
1734 fn execute<'a>(
1735 &'a self,
1736 _ctx: &'a crate::request_context::RequestContext,
1737 _queue: &'a dyn crate::streaming::EventQueueWriter,
1738 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = A2aResult<()>> + Send + 'a>>
1739 {
1740 Box::pin(async { Err(A2aError::internal("streaming fail")) })
1741 }
1742 }
1743
1744 let handler = RequestHandlerBuilder::new(FailExecutor).build().unwrap();
1745 let params = make_params(None);
1746
1747 let result = handler.on_send_message(params, true, None).await;
1748 assert!(
1749 matches!(result, Ok(SendMessageResult::Stream(_))),
1750 "streaming executor failure should still return stream"
1751 );
1752 }
1753
1754 #[tokio::test]
1755 async fn input_required_continuation_reuses_task_id() {
1756 use a2a_protocol_types::task::{Task, TaskId, TaskState, TaskStatus};
1760
1761 let handler = make_handler();
1762
1763 let existing_task_id = TaskId::new("input-required-task");
1765 let task = Task {
1766 id: existing_task_id.clone(),
1767 context_id: ContextId::new("ctx-input"),
1768 status: TaskStatus::new(TaskState::InputRequired),
1769 history: None,
1770 artifacts: None,
1771 metadata: None,
1772 };
1773 handler.task_store.save(&task).await.unwrap();
1774
1775 let mut params = make_params(Some("ctx-input"));
1777 params.message.task_id = Some(existing_task_id.clone());
1778
1779 let result = handler.on_send_message(params, false, None).await;
1780 let send_result = result.expect("continuation should succeed");
1781 match send_result {
1782 SendMessageResult::Response(SendMessageResponse::Task(t)) => {
1783 assert_eq!(
1784 t.id, existing_task_id,
1785 "task_id should be reused for input-required continuation"
1786 );
1787 }
1788 _ => panic!("expected Response(Task)"),
1789 }
1790 }
1791
1792 #[test]
1795 fn second_send_blocked_iff_token_live() {
1796 let live = CancellationEntry {
1797 token: tokio_util::sync::CancellationToken::new(),
1798 created_at: Instant::now(),
1799 };
1800 assert!(
1801 second_send_blocked(&live),
1802 "a live token means an executor is in flight → block the second send"
1803 );
1804
1805 let token = tokio_util::sync::CancellationToken::new();
1806 token.cancel();
1807 let cancelled = CancellationEntry {
1808 token,
1809 created_at: Instant::now(),
1810 };
1811 assert!(
1812 !second_send_blocked(&cancelled),
1813 "a cancelled token no longer blocks a resend"
1814 );
1815 }
1816
1817 #[test]
1818 fn token_aged_at_or_past_max_age() {
1819 let max = std::time::Duration::from_secs(3600);
1820 assert!(
1821 !token_aged(std::time::Duration::from_secs(3599), max),
1822 "younger than max is not aged"
1823 );
1824 assert!(
1827 token_aged(std::time::Duration::from_secs(3600), max),
1828 "exactly max_age is aged"
1829 );
1830 assert!(token_aged(std::time::Duration::from_secs(3601), max));
1831 }
1832
1833 #[test]
1834 fn evict_aged_token_only_when_queue_gone() {
1835 assert!(
1836 evict_aged_token(false),
1837 "no live queue → the executor finished → evict the lingering token"
1838 );
1839 assert!(
1840 !evict_aged_token(true),
1841 "a live queue means the task is still running → keep its token"
1842 );
1843 }
1844
1845 #[test]
1850 fn token_still_evictable_spares_fresh_live_token() {
1851 let max_age = std::time::Duration::from_secs(3600);
1852 let now = Instant::now();
1853
1854 let fresh = CancellationEntry {
1856 token: tokio_util::sync::CancellationToken::new(),
1857 created_at: now,
1858 };
1859 assert!(
1860 !token_still_evictable(&fresh, now, max_age),
1861 "a fresh live token must never be swept"
1862 );
1863
1864 let cancelled = CancellationEntry {
1866 token: tokio_util::sync::CancellationToken::new(),
1867 created_at: now,
1868 };
1869 cancelled.token.cancel();
1870 assert!(token_still_evictable(&cancelled, now, max_age));
1871
1872 let aged = CancellationEntry {
1879 token: tokio_util::sync::CancellationToken::new(),
1880 created_at: now,
1881 };
1882 let later = now
1883 .checked_add(max_age)
1884 .expect("now + max_age is representable");
1885 assert!(token_still_evictable(&aged, later, max_age));
1886 }
1887}