use std::sync::{Arc, Weak};
use ::a2a::{
A2AError, AgentCapabilities, AgentCard, AgentInterface, AgentSkill, Artifact, Message, Part,
PartContent, Role, StreamResponse, TaskArtifactUpdateEvent, TaskState, TaskStatus,
TaskStatusUpdateEvent,
};
use a2a_server::{
AgentExecutor, DefaultRequestHandler, ExecutorContext, InMemoryTaskStore,
WELL_KNOWN_AGENT_CARD_PATH, jsonrpc::jsonrpc_router,
};
use axum::extract::{Path, State};
use axum::http::{HeaderMap, header};
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post};
use axum::{Json, Router};
use futures::stream::BoxStream;
use serde_json::json;
use tokio::sync::mpsc;
use crate::host::{ApiError, Host, NewSession};
pub(crate) fn route(agent: &str) -> String {
format!("/v1/channels/{agent}/a2a")
}
fn thread_channel(agent: &str) -> String {
format!("a2a:{agent}")
}
pub(crate) fn routes(host: &Arc<Host>, mut router: Router<Arc<Host>>) -> Router<Arc<Host>> {
let agents: Vec<String> = host
.app
.inner
.agents
.iter()
.filter(|agent| !agent.sub)
.map(|agent| agent.name.clone())
.collect();
for agent in agents {
let executor = SessionExecutor {
host: Arc::downgrade(host),
agent: agent.clone(),
};
let handler = DefaultRequestHandler::new(executor, InMemoryTaskStore::default())
.with_capabilities(capabilities());
let card = CardSource {
host: Arc::downgrade(host),
agent: agent.clone(),
};
let endpoint = jsonrpc_router(Arc::new(handler))
.route(WELL_KNOWN_AGENT_CARD_PATH, get(agent_card).with_state(card));
router = router
.nest_service(&route(&agent), endpoint.clone())
.nest_service(
&route(&agent).replacen("/v1/channels/", "/v1/e/", 1),
endpoint,
);
}
router
.route("/v1/channels/{name}/a2a", post(unknown))
.route("/v1/e/{name}/a2a", post(unknown))
.route(
&format!("/v1/e/{{name}}/a2a{WELL_KNOWN_AGENT_CARD_PATH}"),
get(unknown),
)
.route(
&format!("/v1/channels/{{name}}/a2a{WELL_KNOWN_AGENT_CARD_PATH}"),
get(unknown),
)
}
async fn unknown(Path(name): Path<String>) -> Response {
crate::server::Failure::from(ApiError::NotFound(format!("agent {name}"))).into_response()
}
fn capabilities() -> AgentCapabilities {
AgentCapabilities {
streaming: Some(true),
push_notifications: Some(false),
extensions: None,
extended_agent_card: None,
}
}
#[derive(Clone)]
struct CardSource {
host: Weak<Host>,
agent: String,
}
async fn agent_card(State(source): State<CardSource>, headers: HeaderMap) -> Response {
let Some(host) = source.host.upgrade() else {
return crate::server::Failure::from(anyhow::anyhow!("host is shutting down"))
.into_response();
};
let url = match headers
.get(header::HOST)
.and_then(|value| value.to_str().ok())
{
Some(authority) => {
let scheme = headers
.get("x-forwarded-proto")
.and_then(|value| value.to_str().ok())
.unwrap_or("http");
format!("{scheme}://{authority}{}", route(&source.agent))
}
None => route(&source.agent),
};
let card = card(&host, &source.agent, url);
([(header::ACCESS_CONTROL_ALLOW_ORIGIN, "*")], Json(card)).into_response()
}
fn card(host: &Host, agent: &str, url: String) -> AgentCard {
let manifest = host.app.manifest();
let description = manifest
.agents
.iter()
.find(|info| info.name == agent)
.and_then(|info| info.description.clone())
.unwrap_or_else(|| format!("The `{agent}` agent of {}", manifest.app.name));
AgentCard {
name: agent.to_string(),
description: description.clone(),
version: manifest.app.version.clone(),
supported_interfaces: vec![AgentInterface::new(url, "JSONRPC")],
capabilities: capabilities(),
default_input_modes: vec!["text/plain".into()],
default_output_modes: vec!["text/plain".into()],
skills: vec![AgentSkill {
id: agent.to_string(),
name: agent.to_string(),
description,
tags: vec!["serve".into(), "everruns".into()],
examples: None,
input_modes: None,
output_modes: None,
security_requirements: None,
}],
provider: None,
documentation_url: None,
icon_url: None,
security_schemes: None,
security_requirements: None,
signatures: None,
}
}
type Events = BoxStream<'static, Result<StreamResponse, A2AError>>;
struct SessionExecutor {
host: Weak<Host>,
agent: String,
}
impl AgentExecutor for SessionExecutor {
fn execute(&self, ctx: ExecutorContext) -> Events {
let (tx, rx) = mpsc::channel(8);
let host = self.host.clone();
let agent = self.agent.clone();
tokio::spawn(async move { run_task(host, agent, ctx, tx).await });
receive(rx)
}
fn cancel(&self, ctx: ExecutorContext) -> Events {
let host = self.host.clone();
let channel = thread_channel(&self.agent);
Box::pin(futures::stream::once(async move {
if let Some(host) = host.upgrade()
&& let Ok(Some(session)) = host.thread_session(&channel, &ctx.context_id)
{
let _ = host.cancel(&session).await;
}
Ok(status(&ctx, TaskState::Canceled, None))
}))
}
}
fn receive(rx: mpsc::Receiver<Result<StreamResponse, A2AError>>) -> Events {
Box::pin(futures::stream::unfold(rx, |mut rx| async move {
rx.recv().await.map(|event| (event, rx))
}))
}
async fn run_task(
host: Weak<Host>,
agent: String,
ctx: ExecutorContext,
tx: mpsc::Sender<Result<StreamResponse, A2AError>>,
) {
let Some(input) = ctx.message.as_ref().and_then(message_input) else {
let _ = tx
.send(Err(A2AError::invalid_params(
"the message needs a text or data part",
)))
.await;
return;
};
let Some(host) = host.upgrade() else {
let _ = tx
.send(Err(A2AError::internal("host is shutting down")))
.await;
return;
};
if tx
.send(Ok(status(&ctx, TaskState::Working, None)))
.await
.is_err()
{
return;
}
let outcome = async {
let session = context_session(&host, &agent, &ctx.context_id).await?;
host.send(&session, input).await?.wait().await
}
.await;
let events = match outcome {
Ok(turn) if turn.success => vec![
StreamResponse::ArtifactUpdate(TaskArtifactUpdateEvent {
task_id: ctx.task_id.clone(),
context_id: ctx.context_id.clone(),
artifact: Artifact {
artifact_id: ::a2a::new_artifact_id(),
name: Some("response".into()),
description: None,
parts: vec![Part::text(turn.response)],
metadata: None,
extensions: None,
},
append: None,
last_chunk: Some(true),
metadata: None,
}),
status(&ctx, TaskState::Completed, None),
],
Ok(turn) => {
let why = turn.error.unwrap_or_else(|| "the turn failed".into());
vec![status(&ctx, TaskState::Failed, Some(why))]
}
Err(err) => vec![status(&ctx, TaskState::Failed, Some(format!("{err:#}")))],
};
for event in events {
if tx.send(Ok(event)).await.is_err() {
return;
}
}
}
async fn context_session(host: &Host, agent: &str, context: &str) -> crate::Result<String> {
let channel = thread_channel(agent);
if let Some(session) = host.thread_session(&channel, context)? {
return Ok(session);
}
let session = host
.create_session(NewSession {
agent: Some(agent.to_string()),
metadata: Some(json!({ "channel": "a2a", "agent": agent, "context": context })),
..NewSession::default()
})
.await?;
host.bind_thread(&channel, context, &session)?;
Ok(session)
}
fn message_input(message: &Message) -> Option<String> {
let parts: Vec<String> = message
.parts
.iter()
.filter_map(|part| match &part.content {
PartContent::Text(text) if !text.trim().is_empty() => Some(text.clone()),
PartContent::Data(value) => Some(value.to_string()),
_ => None,
})
.collect();
(!parts.is_empty()).then(|| parts.join("\n\n"))
}
fn status(ctx: &ExecutorContext, state: TaskState, text: Option<String>) -> StreamResponse {
let message = text.map(|text| {
let mut message = Message::new(Role::Agent, vec![Part::text(text)]);
message.context_id = Some(ctx.context_id.clone());
message.task_id = Some(ctx.task_id.clone());
message
});
StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
task_id: ctx.task_id.clone(),
context_id: ctx.context_id.clone(),
status: TaskStatus {
state,
message,
timestamp: None,
},
metadata: None,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn message_input_joins_text_and_data_parts() {
let message = Message::new(
Role::User,
vec![
Part::text("Summarize"),
Part::text(" "),
Part::data(json!({ "topic": "tides" })),
],
);
assert_eq!(
message_input(&message).unwrap(),
"Summarize\n\n{\"topic\":\"tides\"}"
);
assert!(message_input(&Message::new(Role::User, vec![Part::text(" ")])).is_none());
}
}