mod execution;
mod handlers;
mod message;
mod routes;
use std::collections::{HashMap, HashSet};
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;
use super::agent_adapter::AgentExecutorChain;
use super::protocol::{
A2AErrorData, A2AMessage, A2ARequest, A2AResponse, A2ATask, 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};
use execution::{run_task, run_workflow, sweep_expired_tasks, InflightResume, MAX_WORKFLOW_STEPS};
use handlers::{forbidden, publish_status, task_details_response, task_not_found};
use message::extract_message;
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>>>,
inflight_resumes: Arc<std::sync::Mutex<HashSet<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())),
inflight_resumes: Arc::new(std::sync::Mutex::new(HashSet::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 Err(resp) = self.check_auth(bearer) {
return resp;
}
self.handle_a2a_request(req).await
}
pub(crate) fn check_auth(&self, bearer: Option<&str>) -> Result<(), A2AResponse> {
if let Some(expected) = &self.expected_token {
match bearer {
None => return Err(A2AResponse::error(0, 401, "Authentication required")),
Some(token) if token != expected => {
return Err(A2AResponse::error(0, 401, "Invalid authentication token"));
}
Some(_) => {}
}
}
Ok(())
}
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()),
}
}
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 reserve_message_id(&self, mid: &str) -> Result<Option<String>, ()> {
let mapped = { self.message_ids.read().await.get(mid).cloned() };
if let Some(task_id) = mapped {
if !task_id.is_empty() {
return match self.store.get(&task_id).await {
Ok(Some(_)) => Ok(Some(task_id)),
_ => self.claim_message_id(mid).await,
};
}
return Err(());
}
self.claim_message_id(mid).await
}
async fn claim_message_id(&self, mid: &str) -> Result<Option<String>, ()> {
let mut guard = self.message_ids.write().await;
match guard.get(mid).cloned() {
Some(task_id) if !task_id.is_empty() => Ok(Some(task_id)), Some(_) => Err(()), None => {
guard.insert(mid.to_string(), String::new());
Ok(None)
}
}
}
async fn finish_message_id(&self, mid: &str, task_id: &str) {
self.message_ids
.write()
.await
.insert(mid.to_string(), task_id.to_string());
}
async fn abort_message_id(&self, mid: &str) {
self.message_ids.write().await.remove(mid);
}
async fn release_resume_id(&self, message_id: &Option<String>) {
if let Some(mid) = message_id {
self.abort_message_id(mid).await;
}
}
fn begin_resume(&self, task_id: &str) -> Option<InflightResume> {
let mut guard = self
.inflight_resumes
.lock()
.unwrap_or_else(|e| e.into_inner());
if guard.contains(task_id) {
return None;
}
guard.insert(task_id.to_string());
Some(InflightResume {
inner: self.inflight_resumes.clone(),
task_id: task_id.to_string(),
})
}
async fn cleanup_expired_tasks(&self) {
let Some(ttl) = self.task_ttl else {
return;
};
sweep_expired_tasks(&self.store, ttl).await;
}
}
#[cfg(test)]
mod tests;