use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use serde_json::{json, Value};
use tokio::sync::{broadcast, RwLock};
use lc_agents::AgentExecutor;
use lc_chains::base::{BaseChain, ChainError, ChainResult};
use super::agent_adapter::AgentExecutorChain;
use super::protocol::{
A2AErrorData, A2AMessage, A2ARequest, A2AResponse, A2ATask, A2ATaskResult, A2AWorkflow,
AgentCard, AgentSkill, TaskFilter, TaskPushNotification, TaskStatus,
};
use super::rate_limiter::RateLimiter;
use super::router::{SkillMapRouter, SkillRouter};
use super::store::{InMemoryTaskStore, StoredTask, TaskStore, DEFAULT_MAX_TASKS};
const DEFAULT_TASK_TTL: Duration = Duration::from_secs(24 * 60 * 60);
pub struct A2AServer {
chain: Arc<dyn BaseChain>,
card: AgentCard,
store: Arc<dyn TaskStore>,
message_ids: Arc<RwLock<HashMap<String, String>>>,
skill_router: Option<Arc<dyn SkillRouter>>,
event_bus: Option<Arc<broadcast::Sender<TaskPushNotification>>>,
expected_token: Option<String>,
rate_limiter: Option<Arc<RateLimiter>>,
task_ttl: Option<Duration>,
}
impl A2AServer {
pub fn new(chain: Arc<dyn BaseChain>) -> Self {
let card = AgentCard::new(
chain.name(),
format!("Agent backed by {}", chain.name()),
"http://localhost:8080",
)
.with_skill(AgentSkill::new(
"default",
chain.name(),
format!("Agent backed by {}", chain.name()),
));
Self {
chain,
card,
store: Arc::new(InMemoryTaskStore::with_max_tasks(DEFAULT_MAX_TASKS)),
message_ids: Arc::new(RwLock::new(HashMap::new())),
skill_router: None,
event_bus: None,
expected_token: None,
rate_limiter: None,
task_ttl: Some(DEFAULT_TASK_TTL),
}
}
pub fn from_agent(executor: Arc<AgentExecutor>) -> Self {
Self::new(Arc::new(AgentExecutorChain::new(executor)))
}
pub fn with_store(mut self, store: Arc<dyn TaskStore>) -> Self {
self.store = store;
self
}
pub fn with_max_tasks(mut self, max: usize) -> Self {
self.store = Arc::new(InMemoryTaskStore::with_max_tasks(max.max(1)));
self
}
pub fn with_skill_router(mut self, router: Arc<dyn SkillRouter>) -> Self {
self.skill_router = Some(router);
self
}
pub fn with_skill_map(mut self, map: SkillMapRouter) -> Self {
self.skill_router = Some(Arc::new(map));
self
}
pub fn with_streaming(mut self, capacity: usize) -> Self {
let (tx, _rx) = broadcast::channel(capacity.max(1));
self.event_bus = Some(Arc::new(tx));
self.card = self.card.clone().with_interfaces(json!({ "sse": true }));
self
}
pub fn subscribe(&self) -> Option<broadcast::Receiver<TaskPushNotification>> {
self.event_bus.as_ref().map(|tx| tx.subscribe())
}
pub fn with_auth_token(mut self, token: impl Into<String>) -> Self {
self.expected_token = Some(token.into());
self.card = self
.card
.clone()
.with_authentication(vec!["bearer".to_string()]);
self
}
pub fn with_rate_limiter(mut self, limiter: Arc<RateLimiter>) -> Self {
self.rate_limiter = Some(limiter);
self
}
pub fn with_task_ttl(mut self, ttl: Option<Duration>) -> Self {
self.task_ttl = ttl;
self
}
pub fn with_background_cleanup(self, interval: Duration) -> Self {
let Some(ttl) = self.task_ttl else {
return self;
};
let store = self.store.clone();
let interval = interval.max(Duration::from_millis(1));
tokio::spawn(async move {
let mut ticker = tokio::time::interval(interval);
ticker.tick().await;
loop {
ticker.tick().await;
sweep_expired_tasks(&store, ttl).await;
}
});
self
}
pub fn with_card(mut self, card: AgentCard) -> Self {
self.card = card;
self
}
pub fn get_agent_card(&self) -> &AgentCard {
&self.card
}
pub async fn handle_a2a_request(&self, req: A2ARequest) -> A2AResponse {
if let Some(limiter) = &self.rate_limiter {
if let Err(e) = limiter.try_acquire().await {
return A2AResponse::error(req.id, 429, e.to_string());
}
}
self.dispatch(req).await
}
pub async fn handle_a2a_request_authenticated(
&self,
req: A2ARequest,
bearer: Option<&str>,
) -> A2AResponse {
if let Some(expected) = &self.expected_token {
match bearer {
None => return A2AResponse::error(req.id, 401, "Authentication required"),
Some(token) if token != expected => {
return A2AResponse::error(req.id, 401, "Invalid authentication token");
}
Some(_) => {}
}
}
self.handle_a2a_request(req).await
}
async fn dispatch(&self, req: A2ARequest) -> A2AResponse {
if let Some(trace_id) = req.trace_id() {
log::debug!(
"a2a request method={} id={} trace_id={}",
req.method,
req.id,
trace_id
);
}
match req.method.as_str() {
"tasks/send" => self.handle_tasks_send(req).await,
"tasks/get" => self.handle_tasks_get(req).await,
"tasks/cancel" => self.handle_tasks_cancel(req).await,
"tasks/list" => self.handle_tasks_list(req).await,
"tasks/runWorkflow" => self.handle_workflow_run(req).await,
_ => A2AResponse::from_error_data(req.id, A2AErrorData::method_not_found()),
}
}
async fn handle_tasks_send(&self, req: A2ARequest) -> A2AResponse {
let params = match req.params.clone() {
Some(p) => p,
None => {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::invalid_params("Missing params for tasks/send"),
)
}
};
let message = extract_message(¶ms);
let message_id = req.message_id().map(|s| s.to_string());
if let Some(mid) = &message_id {
if let Some(existing_id) = self.message_ids.read().await.get(mid).cloned() {
if let Ok(Some(stored)) = self.store.get(&existing_id).await {
if !self.caller_owns(&req, &stored.task) {
return forbidden(req.id, "caller does not own the existing task");
}
return A2AResponse::ok(req.id, json!({ "task": stored.task }));
}
self.message_ids.write().await.remove(mid);
}
}
if let Some(task_id) = req.task_id().map(|s| s.to_string()) {
return self
.handle_tasks_send_continue(req, task_id, message, message_id)
.await;
}
let task_id = uuid::Uuid::new_v4().to_string();
let mut task = A2ATask::new(task_id.clone(), message);
if let Some(owner) = req.owner() {
task = task.with_owner(owner);
}
let mut stored = StoredTask::new(task.clone());
if let Some(trace_id) = req.trace_id() {
stored = stored.with_trace_id(trace_id);
}
if self.store.upsert(stored).await.is_err() {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::internal_error("task store write failed"),
);
}
if let Some(mid) = &message_id {
self.message_ids
.write()
.await
.insert(mid.clone(), task_id.clone());
}
let skill_id = params.get("skillId").and_then(Value::as_str);
let chain = self.resolve_chain(skill_id);
let store = self.store.clone();
let bus = self.event_bus.clone();
let history = task.message_history().into_owned();
let spawned_id = task_id.clone();
tokio::spawn(async move {
run_task(&store, chain, &spawned_id, history, bus).await;
});
A2AResponse::ok(req.id, json!({ "task": task }))
}
async fn handle_tasks_send_continue(
&self,
req: A2ARequest,
task_id: String,
message: A2AMessage,
message_id: Option<String>,
) -> A2AResponse {
let mut stored = match self.store.get(&task_id).await {
Ok(Some(s)) => s,
Ok(None) | Err(_) => return task_not_found(req.id, &task_id),
};
if !self.caller_owns(&req, &stored.task) {
return forbidden(req.id, "caller does not own this task");
}
if stored.task.status != TaskStatus::InputRequired {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::new(
-32004,
format!("Cannot continue task in state {}", stored.task.status),
),
);
}
stored.task.push_message(message);
if stored.task.status.can_transition_to(&TaskStatus::Working) {
stored.task.status = TaskStatus::Working;
}
stored.touch();
let task = stored.task.clone();
if self.store.upsert(stored).await.is_err() {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::internal_error("task store write failed"),
);
}
if let Some(mid) = message_id {
self.message_ids.write().await.insert(mid, task_id.clone());
}
let store = self.store.clone();
let chain = self.resolve_chain(None);
let bus = self.event_bus.clone();
let history = task.message_history().into_owned();
let spawned_id = task_id.clone();
tokio::spawn(async move {
run_task(&store, chain, &spawned_id, history, bus).await;
});
A2AResponse::ok(req.id, json!({ "task": task }))
}
async fn handle_tasks_get(&self, req: A2ARequest) -> A2AResponse {
let task_id = req
.params
.as_ref()
.and_then(|p| p.get("taskId"))
.and_then(|v| v.as_str())
.unwrap_or("");
if task_id.is_empty() {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::invalid_params("Missing taskId parameter"),
);
}
self.cleanup_expired_tasks().await;
match self.store.get(task_id).await {
Ok(Some(stored)) => {
if !self.caller_owns(&req, &stored.task) {
return forbidden(req.id, "caller does not own this task");
}
task_details_response(req.id, &stored)
}
Ok(None) => task_not_found(req.id, task_id),
Err(_) => A2AResponse::from_error_data(
req.id,
A2AErrorData::internal_error("task store read failed"),
),
}
}
async fn handle_tasks_cancel(&self, req: A2ARequest) -> A2AResponse {
let task_id = req
.params
.as_ref()
.and_then(|p| p.get("taskId"))
.and_then(|v| v.as_str())
.unwrap_or("");
if task_id.is_empty() {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::invalid_params("Missing taskId parameter"),
);
}
let mut stored = match self.store.get(task_id).await {
Ok(Some(s)) => s,
Ok(None) | Err(_) => return task_not_found(req.id, task_id),
};
if !self.caller_owns(&req, &stored.task) {
return forbidden(req.id, "caller does not own this task");
}
if stored.task.status.is_terminal() {
return A2AResponse::ok(req.id, json!({ "task": stored.task }));
}
if stored.task.status.can_transition_to(&TaskStatus::Cancelled) {
stored.task.status = TaskStatus::Cancelled;
stored.touch();
if self.store.upsert(stored.clone()).await.is_err() {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::internal_error("task store write failed"),
);
}
publish_status(&self.event_bus, task_id, TaskStatus::Cancelled, None);
A2AResponse::ok(req.id, json!({ "task": stored.task }))
} else {
A2AResponse::from_error_data(
req.id,
A2AErrorData::new(
-32002,
format!("Cannot cancel task in state {}", stored.task.status),
),
)
}
}
async fn handle_tasks_list(&self, req: A2ARequest) -> A2AResponse {
self.cleanup_expired_tasks().await;
let mut filter = TaskFilter::new();
if let Some(params) = &req.params {
if let Some(owner) = params.get("owner").and_then(Value::as_str) {
filter = filter.with_owner(owner);
}
if let Some(status) = params.get("status").and_then(Value::as_str) {
if let Ok(ts) =
serde_json::from_value::<TaskStatus>(Value::String(status.to_string()))
{
filter = filter.with_statuses(vec![ts]);
}
}
}
if filter.owner.is_none() {
if let Some(owner) = req.owner() {
filter = filter.with_owner(owner);
}
}
match self.store.list(&filter).await {
Ok(stored) => {
let tasks: Vec<&A2ATask> = stored.iter().map(|s| &s.task).collect();
A2AResponse::ok(req.id, json!({ "tasks": tasks }))
}
Err(_) => A2AResponse::from_error_data(
req.id,
A2AErrorData::internal_error("task store read failed"),
),
}
}
async fn handle_workflow_run(&self, req: A2ARequest) -> A2AResponse {
let params = match req.params.clone() {
Some(p) => p,
None => {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::invalid_params("Missing params for tasks/runWorkflow"),
);
}
};
let workflow: A2AWorkflow = match params.get("workflow") {
Some(w) => match serde_json::from_value(w.clone()) {
Ok(wf) => wf,
Err(_) => {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::invalid_params("Malformed workflow"),
);
}
},
None => {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::invalid_params("Missing workflow for tasks/runWorkflow"),
);
}
};
if workflow.steps.is_empty() {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::invalid_params("Workflow has no steps"),
);
}
let task_id = workflow
.workflow_id
.clone()
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
let mut task = A2ATask::new(
task_id.clone(),
A2AMessage::user(format!(
"workflow: {}",
workflow.name.as_deref().unwrap_or("unnamed")
)),
)
.with_status(TaskStatus::Working);
if let Some(owner) = req.owner() {
task = task.with_owner(owner);
}
let mut stored = StoredTask::new(task.clone());
if let Some(trace_id) = req.trace_id() {
stored = stored.with_trace_id(trace_id);
}
if self.store.upsert(stored).await.is_err() {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::internal_error("task store write failed"),
);
}
publish_status(&self.event_bus, &task_id, TaskStatus::Working, None);
let mut results = serde_json::Map::new();
let mut failure: Option<(String, String)> = None; for step in &workflow.steps {
let chain = self.resolve_chain(step.skill_id.as_deref());
let input = build_chain_input(&step.message.content, chain.as_ref());
match chain.invoke(input).await {
Ok(result) => {
let output = extract_output(&result);
results.insert(step.id.clone(), Value::String(output));
}
Err(e) => {
failure = Some((step.id.clone(), e.to_string()));
break;
}
}
}
let mut finalize = match self.store.get(&task_id).await {
Ok(Some(s)) => s,
_ => {
return A2AResponse::from_error_data(
req.id,
A2AErrorData::internal_error("workflow task vanished"),
);
}
};
let aggregated = results
.values()
.filter_map(|v| v.as_str())
.collect::<Vec<_>>()
.join("\n");
match failure {
Some((step_id, message)) => {
finalize.task.status = TaskStatus::Failed;
finalize.error = Some(format!("step `{step_id}` failed: {message}"));
let error = finalize.error.clone();
let _ = self.store.upsert(finalize.clone()).await;
publish_status(
&self.event_bus,
&task_id,
TaskStatus::Failed,
error.as_deref(),
);
A2AResponse::ok(
req.id,
json!({
"task": finalize.task,
"error": error,
"results": Value::Object(results),
}),
)
}
None => {
finalize.task.status = TaskStatus::Completed;
finalize.result = Some(A2ATaskResult::new(aggregated));
finalize.error = None;
let _ = self.store.upsert(finalize.clone()).await;
publish_status(&self.event_bus, &task_id, TaskStatus::Completed, None);
publish_artifact(&self.event_bus, &task_id, finalize.result.clone().unwrap());
A2AResponse::ok(
req.id,
json!({
"task": finalize.task,
"result": finalize.result,
"results": Value::Object(results),
}),
)
}
}
}
fn caller_owns(&self, req: &A2ARequest, task: &A2ATask) -> bool {
match &task.owner {
Some(task_owner) => req.owner() == Some(task_owner.as_str()),
None => true,
}
}
fn resolve_chain(&self, skill_id: Option<&str>) -> Arc<dyn BaseChain> {
if let Some(sid) = skill_id {
if let Some(router) = &self.skill_router {
if let Some(chain) = router.chain_for(sid) {
return chain;
}
}
}
self.chain.clone()
}
async fn cleanup_expired_tasks(&self) {
let Some(ttl) = self.task_ttl else {
return;
};
sweep_expired_tasks(&self.store, ttl).await;
}
}
async fn sweep_expired_tasks(store: &Arc<dyn TaskStore>, ttl: Duration) {
let Ok(list) = store.list(&TaskFilter::new()).await else {
return;
};
for stored in list {
if stored.age() < ttl {
continue;
}
if stored.task.status.is_terminal() {
let _ = store.delete(&stored.task.id).await;
} else if stored.task.status.can_transition_to(&TaskStatus::Expired) {
let mut expired = stored;
expired.task.status = TaskStatus::Expired;
expired.touch();
let _ = store.upsert(expired).await;
}
}
}
async fn run_task(
store: &Arc<dyn TaskStore>,
chain: Arc<dyn BaseChain>,
task_id: &str,
history: Vec<A2AMessage>,
event_bus: Option<Arc<broadcast::Sender<TaskPushNotification>>>,
) {
if let Ok(Some(mut stored)) = store.get(task_id).await {
if stored.task.status.can_transition_to(&TaskStatus::Working) {
stored.task.status = TaskStatus::Working;
stored.touch();
let _ = store.upsert(stored).await;
publish_status(&event_bus, task_id, TaskStatus::Working, None);
}
}
let input = build_chain_input_from_history(&history, chain.as_ref());
match chain.invoke(input).await {
Ok(result) => {
let output = extract_output(&result);
if let Ok(Some(mut stored)) = store.get(task_id).await {
if stored.task.status.can_transition_to(&TaskStatus::Completed) {
stored.task.status = TaskStatus::Completed;
stored.result = Some(A2ATaskResult::new(output.clone()));
stored.error = None;
stored.touch();
let _ = store.upsert(stored).await;
publish_status(&event_bus, task_id, TaskStatus::Completed, None);
publish_artifact(&event_bus, task_id, A2ATaskResult::new(output));
}
}
}
Err(e) => {
if is_input_required(&e) {
let prompt = e.to_string();
if let Ok(Some(mut stored)) = store.get(task_id).await {
if stored
.task
.status
.can_transition_to(&TaskStatus::InputRequired)
{
stored.task.status = TaskStatus::InputRequired;
stored.error = Some(prompt.clone());
stored.touch();
let _ = store.upsert(stored).await;
publish_status(
&event_bus,
task_id,
TaskStatus::InputRequired,
Some(&prompt),
);
}
}
} else if let Ok(Some(mut stored)) = store.get(task_id).await {
if stored.task.status.can_transition_to(&TaskStatus::Failed) {
stored.task.status = TaskStatus::Failed;
stored.error = Some(e.to_string());
stored.touch();
let _ = store.upsert(stored).await;
publish_status(
&event_bus,
task_id,
TaskStatus::Failed,
Some(&e.to_string()),
);
}
}
}
}
}
fn is_input_required(e: &ChainError) -> bool {
matches!(e, ChainError::MissingInput(_) | ChainError::InputError(_))
}
fn extract_message(params: &Value) -> A2AMessage {
match params.get("message") {
Some(msg_val) => serde_json::from_value(msg_val.clone()).unwrap_or_else(|_| {
A2AMessage::new(
"user",
msg_val
.get("content")
.and_then(|v| v.as_str())
.unwrap_or(""),
)
}),
None => A2AMessage::user(params.to_string()),
}
}
fn build_chain_input_from_history(
history: &[A2AMessage],
chain: &dyn BaseChain,
) -> HashMap<String, Value> {
let content = if history.len() == 1 {
history[0].content.clone()
} else {
history
.iter()
.map(|m| format!("{}: {}", m.role, m.content))
.collect::<Vec<_>>()
.join("\n")
};
build_chain_input(&content, chain)
}
fn build_chain_input(content: &str, chain: &dyn BaseChain) -> HashMap<String, Value> {
let mut map = HashMap::new();
let input_keys = chain.input_keys();
if let Some(first_key) = input_keys.first() {
map.insert(first_key.to_string(), Value::String(content.to_string()));
} else {
map.insert("input".to_string(), Value::String(content.to_string()));
}
map
}
fn extract_output(result: &ChainResult) -> String {
result
.values()
.next()
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string()
}
fn task_details_response(id: u64, stored: &StoredTask) -> A2AResponse {
let mut result = json!({ "task": stored.task });
if let Some(ref task_result) = stored.result {
result["result"] = json!(task_result);
}
if let Some(ref error) = stored.error {
result["error"] = json!(error);
}
A2AResponse::ok(id, result)
}
fn task_not_found(id: u64, task_id: &str) -> A2AResponse {
A2AResponse::from_error_data(
id,
A2AErrorData::new(-32001, format!("Task not found: {}", task_id)),
)
}
fn forbidden(id: u64, message: impl Into<String>) -> A2AResponse {
A2AResponse::from_error_data(id, A2AErrorData::new(-32003, message))
}
fn publish_status(
bus: &Option<Arc<broadcast::Sender<TaskPushNotification>>>,
task_id: &str,
status: TaskStatus,
error: Option<&str>,
) {
if let Some(sender) = bus {
let event = match error {
Some(e) => TaskPushNotification::status_with_error(task_id, status, e),
None => TaskPushNotification::status(task_id, status),
};
let _ = sender.send(event);
}
}
fn publish_artifact(
bus: &Option<Arc<broadcast::Sender<TaskPushNotification>>>,
task_id: &str,
artifact: A2ATaskResult,
) {
if let Some(sender) = bus {
let _ = sender.send(TaskPushNotification::artifact(task_id, artifact));
}
}
#[cfg(test)]
mod tests {
use super::*;
use lc_agents::{AgentError, AgentFinish, AgentOutput, AgentStep, BaseAgent};
use lc_chains::base::{BaseChain, ChainError, ChainResult};
use tokio::sync::Notify;
use crate::protocol::WorkflowStep;
struct EchoChain;
#[async_trait::async_trait]
impl BaseChain for EchoChain {
fn input_keys(&self) -> Vec<&str> {
vec!["input"]
}
fn output_keys(&self) -> Vec<&str> {
vec!["output"]
}
async fn invoke(&self, inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
let input = inputs.get("input").and_then(|v| v.as_str()).unwrap_or("");
let mut result = HashMap::new();
result.insert("output".to_string(), Value::String(input.to_string()));
Ok(result)
}
fn name(&self) -> &str {
"echo-chain"
}
}
struct FailChain;
#[async_trait::async_trait]
impl BaseChain for FailChain {
fn input_keys(&self) -> Vec<&str> {
vec!["input"]
}
fn output_keys(&self) -> Vec<&str> {
vec!["output"]
}
async fn invoke(&self, _inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
Err(ChainError::ExecutionError(
"intentional failure".to_string(),
))
}
fn name(&self) -> &str {
"fail-chain"
}
}
struct BlockingChain {
started: Arc<Notify>,
release: Arc<Notify>,
}
#[async_trait::async_trait]
impl BaseChain for BlockingChain {
fn input_keys(&self) -> Vec<&str> {
vec!["input"]
}
fn output_keys(&self) -> Vec<&str> {
vec!["output"]
}
async fn invoke(&self, _inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
self.started.notify_one();
self.release.notified().await;
let mut result = HashMap::new();
result.insert("output".to_string(), Value::String("done".to_string()));
Ok(result)
}
fn name(&self) -> &str {
"blocking-chain"
}
}
struct InputRequiredChain;
#[async_trait::async_trait]
impl BaseChain for InputRequiredChain {
fn input_keys(&self) -> Vec<&str> {
vec!["input"]
}
fn output_keys(&self) -> Vec<&str> {
vec!["output"]
}
async fn invoke(&self, inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
let input = inputs.get("input").and_then(|v| v.as_str()).unwrap_or("");
if !input.contains("alice") {
return Err(ChainError::MissingInput(
"please provide your name".to_string(),
));
}
let mut result = HashMap::new();
result.insert("output".to_string(), Value::String(input.to_string()));
Ok(result)
}
fn name(&self) -> &str {
"input-required-chain"
}
}
struct NamedChain(String);
#[async_trait::async_trait]
impl BaseChain for NamedChain {
fn input_keys(&self) -> Vec<&str> {
vec!["input"]
}
fn output_keys(&self) -> Vec<&str> {
vec!["output"]
}
async fn invoke(&self, _inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
let mut out = HashMap::new();
out.insert("output".to_string(), Value::String(self.0.clone()));
Ok(out)
}
fn name(&self) -> &str {
&self.0
}
}
fn echo_server() -> A2AServer {
A2AServer::new(Arc::new(EchoChain))
}
fn fail_server() -> A2AServer {
A2AServer::new(Arc::new(FailChain))
}
struct EchoAgent;
#[async_trait::async_trait]
impl BaseAgent for EchoAgent {
async fn plan(
&self,
_intermediate_steps: &[AgentStep],
inputs: &HashMap<String, String>,
) -> Result<AgentOutput, AgentError> {
let input = inputs.get("input").cloned().unwrap_or_default();
Ok(AgentOutput::Finish(AgentFinish::new(
format!("agent-said: {}", input),
String::new(),
)))
}
}
async fn wait_for_status(server: &A2AServer, task_id: &str, want: &str) -> A2AResponse {
for _ in 0..200 {
let resp = server
.handle_a2a_request(A2ARequest::get_task(99, task_id))
.await;
if let Some(r) = &resp.result {
if r.get("task")
.and_then(|t| t.get("status"))
.and_then(|s| s.as_str())
== Some(want)
{
return resp;
}
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
panic!("task {task_id} did not reach status {want} in time");
}
async fn send_task_id(server: &A2AServer, content: &str) -> String {
let resp = server
.handle_a2a_request(A2ARequest::send_task(1, &A2AMessage::user(content)))
.await;
assert!(!resp.is_error());
resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string()
}
#[test]
fn get_agent_card_default() {
let server = echo_server();
let card = server.get_agent_card();
assert_eq!(card.name, "echo-chain");
assert!(card.description.contains("echo-chain"));
assert_eq!(card.protocol_version, "0.3.0");
assert_eq!(card.skills.len(), 1);
assert_eq!(card.skills[0].id, "default");
assert_eq!(card.skills[0].name, "echo-chain");
}
#[test]
fn get_agent_card_custom() {
let card = AgentCard::new("custom", "Custom agent", "http://example.com")
.with_skill(AgentSkill::new("s1", "text-generation", "Generates text"));
let server = echo_server().with_card(card);
let card = server.get_agent_card();
assert_eq!(card.name, "custom");
assert_eq!(card.url, "http://example.com");
assert_eq!(card.skills.len(), 1);
assert_eq!(card.skills[0].id, "s1");
}
#[tokio::test]
async fn handle_tasks_send_returns_submitted_immediately() {
let server = echo_server();
let msg = A2AMessage::user("hello world");
let req = A2ARequest::send_task(1, &msg);
let resp = server.handle_a2a_request(req).await;
assert!(!resp.is_error());
let result = resp.result.unwrap();
let task = result.get("task").unwrap();
assert_eq!(task["status"], "submitted");
assert!(result.get("result").is_none());
}
#[tokio::test]
async fn from_agent_serves_tasks_end_to_end() {
let executor = Arc::new(AgentExecutor::new(Arc::new(EchoAgent), Vec::new()));
let server = A2AServer::from_agent(executor);
let task_id = send_task_id(&server, "hi").await;
let done = wait_for_status(&server, &task_id, "completed").await;
let output = done.result.unwrap()["result"]["output"]
.as_str()
.unwrap()
.to_string();
assert_eq!(output, "agent-said: hi");
}
#[tokio::test]
async fn handle_tasks_get_shows_working_then_completed() {
let server = echo_server();
let msg = A2AMessage::user("hello");
let send_resp = server
.handle_a2a_request(A2ARequest::send_task(1, &msg))
.await;
let task_id = send_resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let done = wait_for_status(&server, &task_id, "completed").await;
let result = done.result.unwrap();
let task = result.get("task").unwrap();
assert_eq!(task["id"], task_id);
assert_eq!(task["status"], "completed");
let task_result = result.get("result").unwrap();
assert_eq!(task_result["output"], "hello");
}
#[tokio::test]
async fn handle_tasks_send_failure() {
let server = fail_server();
let send_resp = server
.handle_a2a_request(A2ARequest::send_task(2, &A2AMessage::user("hello")))
.await;
assert!(!send_resp.is_error());
let task_id = send_resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let done = wait_for_status(&server, &task_id, "failed").await;
let result = done.result.unwrap();
assert_eq!(result["task"]["status"], "failed");
let error = result.get("error").unwrap();
assert!(error.as_str().unwrap().contains("intentional failure"));
}
#[tokio::test]
async fn handle_tasks_send_missing_params() {
let server = echo_server();
let req = A2ARequest::new(3, "tasks/send", None);
let resp = server.handle_a2a_request(req).await;
assert!(resp.is_error());
let err = resp.error.unwrap();
assert_eq!(err.code, -32602);
}
#[tokio::test]
async fn handle_tasks_get_not_found() {
let server = echo_server();
let req = A2ARequest::get_task(5, "nonexistent-task");
let resp = server.handle_a2a_request(req).await;
assert!(resp.is_error());
let err = resp.error.unwrap();
assert!(err.message.contains("Task not found"));
}
#[tokio::test]
async fn handle_tasks_cancel_nonexistent() {
let server = echo_server();
let req = A2ARequest::cancel_task(6, "task-123");
let resp = server.handle_a2a_request(req).await;
assert!(resp.is_error());
let err = resp.error.unwrap();
assert!(err.message.contains("Task not found"));
}
#[tokio::test]
async fn handle_tasks_cancel_working_task() {
let started = Arc::new(Notify::new());
let release = Arc::new(Notify::new());
let chain = Arc::new(BlockingChain {
started: started.clone(),
release: release.clone(),
});
let server = A2AServer::new(chain);
let send_resp = server
.handle_a2a_request(A2ARequest::send_task(1, &A2AMessage::user("hi")))
.await;
let task_id = send_resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
started.notified().await;
let get_resp = server
.handle_a2a_request(A2ARequest::get_task(2, &task_id))
.await;
assert_eq!(get_resp.result.unwrap()["task"]["status"], "working");
let cancel_resp = server
.handle_a2a_request(A2ARequest::cancel_task(3, &task_id))
.await;
assert!(!cancel_resp.is_error());
assert_eq!(cancel_resp.result.unwrap()["task"]["status"], "cancelled");
release.notify_one();
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
let get_resp = server
.handle_a2a_request(A2ARequest::get_task(4, &task_id))
.await;
assert_eq!(get_resp.result.unwrap()["task"]["status"], "cancelled");
}
#[tokio::test]
async fn handle_tasks_cancel_completed_is_idempotent() {
let server = echo_server();
let send_resp = server
.handle_a2a_request(A2ARequest::send_task(1, &A2AMessage::user("hi")))
.await;
let task_id = send_resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
wait_for_status(&server, &task_id, "completed").await;
let cancel_resp = server
.handle_a2a_request(A2ARequest::cancel_task(2, &task_id))
.await;
assert!(!cancel_resp.is_error());
assert_eq!(cancel_resp.result.unwrap()["task"]["status"], "completed");
}
#[tokio::test]
async fn handle_tasks_cancel_missing_task_id() {
let server = echo_server();
let req = A2ARequest::new(7, "tasks/cancel", Some(json!({})));
let resp = server.handle_a2a_request(req).await;
assert!(resp.is_error());
}
#[tokio::test]
async fn handle_tasks_list_returns_tasks() {
let server = echo_server();
let send_resp = server
.handle_a2a_request(A2ARequest::send_task(1, &A2AMessage::user("hi")))
.await;
let task_id = send_resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let resp = server
.handle_a2a_request(A2ARequest::new(2, "tasks/list", None))
.await;
assert!(!resp.is_error());
let result = resp.result.unwrap();
let tasks = result["tasks"].as_array().unwrap();
assert!(
tasks
.iter()
.any(|t| t["id"].as_str() == Some(task_id.as_str())),
"expected task {task_id} in list"
);
}
#[tokio::test]
async fn handle_unknown_method() {
let server = echo_server();
let req = A2ARequest::new(8, "foo/bar", None);
let resp = server.handle_a2a_request(req).await;
assert!(resp.is_error());
let err = resp.error.unwrap();
assert_eq!(err.code, -32601);
}
#[tokio::test]
async fn handle_tasks_send_with_raw_params() {
let server = echo_server();
let req = A2ARequest::new(9, "tasks/send", Some(json!({"query": "test query"})));
let resp = server.handle_a2a_request(req).await;
assert!(!resp.is_error());
}
#[tokio::test]
async fn handle_tasks_send_chain_with_no_input_keys() {
struct NoKeyChain;
#[async_trait::async_trait]
impl BaseChain for NoKeyChain {
fn input_keys(&self) -> Vec<&str> {
vec![]
}
fn output_keys(&self) -> Vec<&str> {
vec!["output"]
}
async fn invoke(
&self,
inputs: HashMap<String, Value>,
) -> Result<ChainResult, ChainError> {
let input = inputs
.get("input")
.and_then(|v| v.as_str())
.unwrap_or("default");
let mut result = HashMap::new();
result.insert("output".to_string(), Value::String(input.to_string()));
Ok(result)
}
fn name(&self) -> &str {
"no-key-chain"
}
}
let server = A2AServer::new(Arc::new(NoKeyChain));
let msg = A2AMessage::user("hello");
let req = A2ARequest::send_task(10, &msg);
let resp = server.handle_a2a_request(req).await;
assert!(!resp.is_error());
}
#[tokio::test]
async fn handle_a2a_request_authenticated_requires_token() {
let server = echo_server().with_auth_token("secret");
assert_eq!(
server.get_agent_card().authentication,
Some(vec!["bearer".to_string()])
);
let msg = A2AMessage::user("hi");
let resp = server
.handle_a2a_request_authenticated(A2ARequest::send_task(1, &msg), None)
.await;
assert!(resp.is_error());
assert_eq!(resp.error.unwrap().code, 401);
}
#[tokio::test]
async fn handle_a2a_request_authenticated_invalid_token() {
let server = echo_server().with_auth_token("secret");
let msg = A2AMessage::user("hi");
let resp = server
.handle_a2a_request_authenticated(A2ARequest::send_task(1, &msg), Some("wrong"))
.await;
assert!(resp.is_error());
assert_eq!(resp.error.unwrap().code, 401);
}
#[tokio::test]
async fn handle_a2a_request_authenticated_valid_token() {
let server = echo_server().with_auth_token("secret");
let msg = A2AMessage::user("hi");
let resp = server
.handle_a2a_request_authenticated(A2ARequest::send_task(1, &msg), Some("secret"))
.await;
assert!(!resp.is_error());
}
#[tokio::test]
async fn handle_a2a_request_unauthenticated_passes() {
let server = echo_server();
let msg = A2AMessage::user("hi");
let resp = server
.handle_a2a_request_authenticated(A2ARequest::send_task(1, &msg), None)
.await;
assert!(!resp.is_error());
}
#[tokio::test]
async fn handle_a2a_request_rate_limited() {
let server = echo_server().with_rate_limiter(Arc::new(RateLimiter::new(0, 1)));
let msg = A2AMessage::user("hi");
let r1 = server
.handle_a2a_request(A2ARequest::send_task(1, &msg))
.await;
assert!(!r1.is_error());
let r2 = server
.handle_a2a_request(A2ARequest::send_task(2, &msg))
.await;
assert!(r2.is_error());
assert_eq!(r2.error.unwrap().code, 429);
}
#[tokio::test]
async fn handle_tasks_send_idempotent_message_id() {
let server = echo_server();
let msg = A2AMessage::user("hello");
let req = A2ARequest::send_task_with_message_id(1, &msg, "idem-1");
let r1 = server.handle_a2a_request(req.clone()).await;
assert!(!r1.is_error());
let task1 = r1.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let r2 = server.handle_a2a_request(req).await;
assert!(!r2.is_error());
let task2 = r2.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
assert_eq!(task1, task2);
}
#[tokio::test]
async fn handle_tasks_get_owner_enforced() {
let server = echo_server();
let send_resp = server
.handle_a2a_request(
A2ARequest::send_task(1, &A2AMessage::user("hi")).with_owner("alice"),
)
.await;
let task_id = send_resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let ok = server
.handle_a2a_request(A2ARequest::get_task(2, &task_id).with_owner("alice"))
.await;
assert!(!ok.is_error());
let denied = server
.handle_a2a_request(A2ARequest::get_task(3, &task_id).with_owner("bob"))
.await;
assert!(denied.is_error());
assert_eq!(denied.error.unwrap().code, -32003);
let anon = server
.handle_a2a_request(A2ARequest::get_task(4, &task_id))
.await;
assert!(anon.is_error());
assert_eq!(anon.error.unwrap().code, -32003);
}
#[tokio::test]
async fn handle_tasks_cancel_owner_enforced() {
let server = echo_server();
let send_resp = server
.handle_a2a_request(
A2ARequest::send_task(1, &A2AMessage::user("hi")).with_owner("alice"),
)
.await;
let task_id = send_resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let denied = server
.handle_a2a_request(A2ARequest::cancel_task(2, &task_id).with_owner("bob"))
.await;
assert!(denied.is_error());
assert_eq!(denied.error.unwrap().code, -32003);
let ok = server
.handle_a2a_request(A2ARequest::cancel_task(3, &task_id).with_owner("alice"))
.await;
assert!(!ok.is_error());
assert_eq!(ok.result.unwrap()["task"]["status"], "cancelled");
}
#[tokio::test]
async fn handle_tasks_list_filters_by_owner() {
let server = echo_server();
let _a = server
.handle_a2a_request(
A2ARequest::send_task(1, &A2AMessage::user("hi")).with_owner("alice"),
)
.await;
let _b = server
.handle_a2a_request(A2ARequest::send_task(2, &A2AMessage::user("hi")).with_owner("bob"))
.await;
let resp = server
.handle_a2a_request(A2ARequest::new(3, "tasks/list", None).with_owner("alice"))
.await;
let tasks = resp.result.unwrap()["tasks"].as_array().unwrap().clone();
assert_eq!(tasks.len(), 1);
assert_eq!(tasks[0]["owner"], "alice");
let resp = server
.handle_a2a_request(A2ARequest::new(4, "tasks/list", None))
.await;
let tasks = resp.result.unwrap()["tasks"].as_array().unwrap().clone();
assert_eq!(tasks.len(), 2);
}
#[tokio::test]
async fn handle_tasks_send_continue_terminal_rejected() {
let server = echo_server();
let task_id = send_task_id(&server, "hello").await;
wait_for_status(&server, &task_id, "completed").await;
let continue_resp = server
.handle_a2a_request(A2ARequest::continue_task(
2,
&task_id,
&A2AMessage::user("x"),
))
.await;
assert!(continue_resp.is_error());
assert_eq!(continue_resp.error.unwrap().code, -32004);
}
#[tokio::test]
async fn handle_tasks_send_input_required_then_resume() {
let server = A2AServer::new(Arc::new(InputRequiredChain));
let task_id = send_task_id(&server, "hello").await;
let pending = wait_for_status(&server, &task_id, "input-required").await;
let error = pending.result.unwrap()["error"]
.as_str()
.unwrap()
.to_string();
assert!(error.contains("please provide your name"));
let resume_resp = server
.handle_a2a_request(A2ARequest::continue_task(
2,
&task_id,
&A2AMessage::user("my name is alice"),
))
.await;
assert!(!resume_resp.is_error());
let done = wait_for_status(&server, &task_id, "completed").await;
let output = done.result.unwrap()["result"]["output"]
.as_str()
.unwrap()
.to_string();
assert!(
output.contains("hello"),
"resumed output missing first turn: {output}"
);
assert!(
output.contains("alice"),
"resumed output missing answer: {output}"
);
}
#[tokio::test]
async fn handle_tasks_send_continue_working_rejected() {
let started = Arc::new(Notify::new());
let release = Arc::new(Notify::new());
let chain = Arc::new(BlockingChain {
started: started.clone(),
release: release.clone(),
});
let server = A2AServer::new(chain);
let send_resp = server
.handle_a2a_request(A2ARequest::send_task(1, &A2AMessage::user("hi")))
.await;
let task_id = send_resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
started.notified().await;
let continue_resp = server
.handle_a2a_request(A2ARequest::continue_task(
2,
&task_id,
&A2AMessage::user("more"),
))
.await;
assert!(continue_resp.is_error());
assert_eq!(continue_resp.error.unwrap().code, -32004);
release.notify_one();
}
#[tokio::test]
async fn with_store_custom_backend() {
let store = InMemoryTaskStore::with_max_tasks(1);
let server = echo_server().with_store(Arc::new(store));
let first = send_task_id(&server, "one").await;
let second = send_task_id(&server, "two").await;
let gone = server
.handle_a2a_request(A2ARequest::get_task(2, &first))
.await;
assert!(gone.is_error());
assert!(gone.error.unwrap().message.contains("Task not found"));
let present = server
.handle_a2a_request(A2ARequest::get_task(3, &second))
.await;
assert!(!present.is_error());
}
#[tokio::test]
async fn handle_tasks_send_routes_by_skill() {
let router = SkillMapRouter::new()
.with_skill("math", Arc::new(NamedChain("math-chain".to_string())));
let server = echo_server().with_skill_map(router);
let params = json!({
"message": { "role": "user", "content": "hi" },
"skillId": "math"
});
let resp = server
.handle_a2a_request(A2ARequest::new(1, "tasks/send", Some(params)))
.await;
let task_id = resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let done = wait_for_status(&server, &task_id, "completed").await;
let output = done.result.unwrap()["result"]["output"]
.as_str()
.unwrap()
.to_string();
assert_eq!(output, "math-chain");
let task_id = send_task_id(&server, "hello").await;
let done = wait_for_status(&server, &task_id, "completed").await;
let output = done.result.unwrap()["result"]["output"]
.as_str()
.unwrap()
.to_string();
assert_eq!(output, "hello");
}
#[tokio::test]
async fn handle_tasks_send_unknown_skill_falls_back() {
let router = SkillMapRouter::new()
.with_skill("math", Arc::new(NamedChain("math-chain".to_string())));
let server = echo_server().with_skill_map(router);
let params = json!({
"message": { "role": "user", "content": "hi" },
"skillId": "unknown"
});
let resp = server
.handle_a2a_request(A2ARequest::new(1, "tasks/send", Some(params)))
.await;
let task_id = resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let done = wait_for_status(&server, &task_id, "completed").await;
let output = done.result.unwrap()["result"]["output"]
.as_str()
.unwrap()
.to_string();
assert_eq!(output, "hi");
}
#[tokio::test]
async fn with_streaming_publishes_events_and_advertises_sse() {
let server = echo_server().with_streaming(16);
let card = server.get_agent_card();
assert_eq!(
card.interfaces,
Some(json!({ "sse": true })),
"streaming advertises sse interface"
);
let mut rx = server.subscribe().expect("subscribed");
let task_id = send_task_id(&server, "hello").await;
let mut saw_working = false;
let mut saw_completed = false;
let mut saw_artifact = false;
for _ in 0..8 {
match rx.recv().await {
Ok(event) => {
assert_eq!(event.id(), task_id);
match event.status_value() {
Some(TaskStatus::Working) => saw_working = true,
Some(TaskStatus::Completed) => saw_completed = true,
None => saw_artifact = true,
_ => {}
}
if saw_completed && saw_artifact {
break;
}
}
Err(_) => break,
}
}
assert!(saw_working, "expected a working event");
assert!(saw_completed, "expected a completed event");
assert!(saw_artifact, "expected an artifact event");
}
#[tokio::test]
async fn without_streaming_no_subscriber() {
let server = echo_server();
assert!(server.subscribe().is_none());
}
#[tokio::test]
async fn sweep_expired_tasks_cleans_terminal_and_expires_live() {
let store: Arc<dyn TaskStore> = Arc::new(InMemoryTaskStore::with_max_tasks(10));
let mut terminal = StoredTask::new(
A2ATask::new("t-term", A2AMessage::user("done")).with_status(TaskStatus::Completed),
);
terminal.updated_at = std::time::Instant::now() - Duration::from_secs(100);
let mut live = StoredTask::new(
A2ATask::new("t-live", A2AMessage::user("run")).with_status(TaskStatus::Working),
);
live.updated_at = std::time::Instant::now() - Duration::from_secs(100);
let fresh = StoredTask::new(A2ATask::new("t-fresh", A2AMessage::user("new")));
store.upsert(terminal).await.unwrap();
store.upsert(live).await.unwrap();
store.upsert(fresh).await.unwrap();
sweep_expired_tasks(&store, Duration::from_secs(10)).await;
assert!(
store.get("t-term").await.unwrap().is_none(),
"terminal task past TTL is deleted"
);
let live = store.get("t-live").await.unwrap().expect("live task kept");
assert_eq!(live.task.status, TaskStatus::Expired);
let fresh = store
.get("t-fresh")
.await
.unwrap()
.expect("fresh task kept");
assert_eq!(fresh.task.status, TaskStatus::Submitted);
}
#[tokio::test]
async fn background_cleanup_sweeps_periodically() {
let store: Arc<dyn TaskStore> = Arc::new(InMemoryTaskStore::with_max_tasks(10));
let mut expired = StoredTask::new(
A2ATask::new("t-old", A2AMessage::user("hi")).with_status(TaskStatus::Working),
);
expired.updated_at = std::time::Instant::now() - Duration::from_secs(100);
store.upsert(expired).await.unwrap();
let _server = A2AServer::new(Arc::new(EchoChain))
.with_store(store.clone())
.with_task_ttl(Some(Duration::from_secs(10)))
.with_background_cleanup(Duration::from_millis(5));
tokio::time::sleep(Duration::from_millis(50)).await;
let t = store
.get("t-old")
.await
.unwrap()
.expect("live task still present");
assert_eq!(t.task.status, TaskStatus::Expired);
}
#[tokio::test]
async fn task_stores_request_trace_id() {
let store: Arc<dyn TaskStore> = Arc::new(InMemoryTaskStore::with_max_tasks(10));
let server = A2AServer::new(Arc::new(EchoChain)).with_store(store.clone());
let req = A2ARequest::send_task(1, &A2AMessage::user("hello")).with_trace_id("trace-abc");
let resp = server.handle_a2a_request(req).await;
assert!(!resp.is_error());
let task_id = resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let stored = store.get(&task_id).await.unwrap().expect("task stored");
assert_eq!(stored.trace_id.as_deref(), Some("trace-abc"));
}
#[tokio::test]
async fn task_without_trace_id_has_none() {
let store: Arc<dyn TaskStore> = Arc::new(InMemoryTaskStore::with_max_tasks(10));
let server = A2AServer::new(Arc::new(EchoChain)).with_store(store.clone());
let resp = server
.handle_a2a_request(A2ARequest::send_task(1, &A2AMessage::user("hi")))
.await;
assert!(!resp.is_error());
let task_id = resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let stored = store.get(&task_id).await.unwrap().expect("task stored");
assert!(stored.trace_id.is_none());
}
#[tokio::test]
async fn run_workflow_executes_steps_in_order_and_aggregates() {
let store: Arc<dyn TaskStore> = Arc::new(InMemoryTaskStore::with_max_tasks(10));
let server = A2AServer::new(Arc::new(EchoChain)).with_store(store.clone());
let workflow = A2AWorkflow::new(vec![
WorkflowStep::new("s1", "first"),
WorkflowStep::new("s2", "second"),
]);
let resp = server
.handle_a2a_request(A2ARequest::run_workflow(1, &workflow))
.await;
assert!(!resp.is_error());
let result = resp.result.unwrap();
assert_eq!(result["task"]["status"], "completed");
assert_eq!(result["results"]["s1"], "first");
assert_eq!(result["results"]["s2"], "second");
let task_id = result["task"]["id"].as_str().unwrap();
let stored = store.get(task_id).await.unwrap().expect("task stored");
assert_eq!(stored.task.status, TaskStatus::Completed);
assert_eq!(stored.result.as_ref().unwrap().output, "first\nsecond");
}
#[tokio::test]
async fn run_workflow_respects_supplied_workflow_id_and_owner() {
let store: Arc<dyn TaskStore> = Arc::new(InMemoryTaskStore::with_max_tasks(10));
let server = A2AServer::new(Arc::new(EchoChain)).with_store(store.clone());
let workflow = A2AWorkflow::new(vec![WorkflowStep::new("s1", "hi")])
.with_workflow_id("wf-42")
.with_name("my workflow");
let resp = server
.handle_a2a_request(A2ARequest::run_workflow(1, &workflow).with_owner("alice"))
.await;
assert!(!resp.is_error());
let result = resp.result.unwrap();
assert_eq!(result["task"]["id"], "wf-42");
let stored = store.get("wf-42").await.unwrap().expect("task stored");
assert_eq!(stored.task.owner.as_deref(), Some("alice"));
}
#[tokio::test]
async fn run_workflow_routes_steps_by_skill() {
let store: Arc<dyn TaskStore> = Arc::new(InMemoryTaskStore::with_max_tasks(10));
let server = A2AServer::new(Arc::new(EchoChain))
.with_store(store.clone())
.with_skill_map(
SkillMapRouter::new()
.with_skill("translate", Arc::new(NamedChain("translated".to_string()))),
);
let workflow = A2AWorkflow::new(vec![
WorkflowStep::new("s1", "hello"),
WorkflowStep::with_skill("s2", "bonjour", "translate"),
]);
let resp = server
.handle_a2a_request(A2ARequest::run_workflow(1, &workflow))
.await;
assert!(!resp.is_error());
let result = resp.result.unwrap();
assert_eq!(result["results"]["s1"], "hello");
assert_eq!(result["results"]["s2"], "translated");
}
#[tokio::test]
async fn run_workflow_step_failure_marks_task_failed_and_stops() {
let store: Arc<dyn TaskStore> = Arc::new(InMemoryTaskStore::with_max_tasks(10));
let failing = A2AServer::new(Arc::new(EchoChain))
.with_store(store.clone())
.with_skill_map(SkillMapRouter::new().with_skill("failing", Arc::new(FailChain)));
let workflow = A2AWorkflow::new(vec![
WorkflowStep::new("s1", "ok"),
WorkflowStep::with_skill("s2", "boom", "failing"),
]);
let resp = failing
.handle_a2a_request(A2ARequest::run_workflow(1, &workflow))
.await;
assert!(!resp.is_error()); let result = resp.result.unwrap();
assert_eq!(result["task"]["status"], "failed");
assert!(result["results"].get("s1").is_some());
assert!(result["results"].get("s2").is_none());
assert!(result["error"]
.as_str()
.unwrap()
.contains("step `s2` failed"));
let stored = store
.get(result["task"]["id"].as_str().unwrap())
.await
.unwrap()
.unwrap();
assert_eq!(stored.task.status, TaskStatus::Failed);
assert!(stored.error.as_deref().unwrap().contains("s2"));
}
#[tokio::test]
async fn run_workflow_missing_params_invalid() {
let server = A2AServer::new(Arc::new(EchoChain));
let resp = server
.handle_a2a_request(A2ARequest::new(1, "tasks/runWorkflow", None))
.await;
assert!(resp.is_error());
}
#[tokio::test]
async fn run_workflow_empty_steps_invalid() {
let server = A2AServer::new(Arc::new(EchoChain));
let resp = server
.handle_a2a_request(A2ARequest::run_workflow(1, &A2AWorkflow::new(vec![])))
.await;
assert!(resp.is_error());
}
#[tokio::test]
async fn run_workflow_carries_trace_id_onto_backing_task() {
let store: Arc<dyn TaskStore> = Arc::new(InMemoryTaskStore::with_max_tasks(10));
let server = A2AServer::new(Arc::new(EchoChain)).with_store(store.clone());
let workflow = A2AWorkflow::new(vec![WorkflowStep::new("s1", "hi")]);
let resp = server
.handle_a2a_request(A2ARequest::run_workflow(1, &workflow).with_trace_id("trace-wf"))
.await;
assert!(!resp.is_error());
let task_id = resp.result.unwrap()["task"]["id"]
.as_str()
.unwrap()
.to_string();
let stored = store.get(&task_id).await.unwrap().unwrap();
assert_eq!(stored.trace_id.as_deref(), Some("trace-wf"));
}
}