Skip to main content

systemprompt_agent/services/a2a_server/
server.rs

1//! The per-agent A2A HTTP server.
2//!
3//! [`Server`] loads an agent's configuration, wires OAuth state and the AI
4//! provider, and builds the axum [`Router`] exposing the agent card and the A2A
5//! request endpoint, then runs the listener.
6//!
7//! Copyright (c) systemprompt.io — Business Source License 1.1.
8//! See <https://systemprompt.io> for licensing details.
9
10use axum::extract::DefaultBodyLimit;
11use axum::http::{HeaderValue, Method};
12use axum::routing::{get, post};
13use axum::{Router, middleware};
14use std::sync::Arc;
15use systemprompt_database::DbPool;
16use systemprompt_models::modules::ApiPaths;
17use systemprompt_models::{AgentConfig, AiProvider};
18use tokio::sync::{RwLock, Semaphore};
19use tower_http::cors::{AllowOrigin, CorsLayer};
20
21use super::active_tasks::ActiveTasks;
22use super::auth::{AgentOAuthConfig, AgentOAuthState, agent_oauth_middleware_wrapper};
23use super::handlers::{AgentHandlerState, handle_agent_card, handle_agent_request};
24use crate::state::AgentState;
25
26// Why: A2A file parts travel inline as base64, so the JSON-RPC body cap is
27// deliberately wider than the API's default request limit.
28pub const A2A_MAX_REQUEST_BODY_BYTES: usize = 8 * 1024 * 1024;
29
30fn cors_layer(origins: &[String]) -> Result<CorsLayer, crate::error::AgentError> {
31    let mut allowed = Vec::new();
32    for origin in origins {
33        let trimmed = origin.trim();
34        if trimmed.is_empty() {
35            continue;
36        }
37        let value = trimmed.parse::<HeaderValue>().map_err(|e| {
38            crate::error::AgentError::Config(format!(
39                "invalid cors_allowed_origins entry {origin:?}: {e}"
40            ))
41        })?;
42        allowed.push(value);
43    }
44    if allowed.is_empty() {
45        return Err(crate::error::AgentError::Config(
46            "cors_allowed_origins must contain at least one valid origin".to_owned(),
47        ));
48    }
49    Ok(CorsLayer::new()
50        .allow_origin(AllowOrigin::list(allowed))
51        .allow_credentials(true)
52        .allow_methods([Method::GET, Method::POST, Method::OPTIONS])
53        .allow_headers([
54            http::header::AUTHORIZATION,
55            http::header::CONTENT_TYPE,
56            http::header::ACCEPT,
57        ]))
58}
59
60pub struct Server {
61    config: Arc<RwLock<AgentConfig>>,
62    oauth_state: Arc<AgentOAuthState>,
63    agent_state: Arc<AgentState>,
64    ai_service: Arc<dyn AiProvider>,
65    stream_semaphore: Arc<Semaphore>,
66    active_tasks: ActiveTasks,
67    cors: CorsLayer,
68    port: u16,
69}
70
71impl std::fmt::Debug for Server {
72    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
73        f.debug_struct("Server")
74            .field("config", &"Arc<RwLock<AgentConfig>>")
75            .field("oauth_state", &"Arc<AgentOAuthState>")
76            .field("agent_state", &"Arc<AgentState>")
77            .field("ai_service", &"<Arc<dyn AiProvider>>")
78            .field(
79                "stream_semaphore",
80                &self.stream_semaphore.available_permits(),
81            )
82            .field("active_tasks", &self.active_tasks)
83            .field("port", &self.port)
84            .finish_non_exhaustive()
85    }
86}
87
88impl Server {
89    pub async fn new(
90        db_pool: DbPool,
91        agent_state: Arc<AgentState>,
92        ai_service: Arc<dyn AiProvider>,
93        agent_name: Option<String>,
94        port: u16,
95    ) -> Result<Self, crate::error::AgentError> {
96        use crate::services::registry::AgentRegistry;
97
98        let mut config = if let Some(name) = agent_name {
99            let registry = AgentRegistry::new()
100                .map_err(|e| crate::error::AgentError::Server(e.to_string()))?;
101            registry
102                .get_agent(&name)
103                .await
104                .map_err(|e| crate::error::AgentError::Server(e.to_string()))?
105        } else {
106            return Err(crate::error::AgentError::Validation(
107                "Agent name is required".to_owned(),
108            ));
109        };
110
111        config.extract_oauth_scopes_from_card();
112
113        let oauth_config = AgentOAuthConfig::default();
114        let global_config = agent_state.config();
115        let mut oauth_state = AgentOAuthState::new(
116            Arc::clone(&db_pool),
117            oauth_config,
118            global_config.jwt_issuer.clone(),
119            global_config.jwt_audiences.clone(),
120        );
121
122        oauth_state = oauth_state.with_jwt_provider(Arc::clone(agent_state.jwt_provider()));
123        let cors = cors_layer(&global_config.cors_allowed_origins)?;
124        let stream_semaphore = Arc::new(Semaphore::new(global_config.max_concurrent_streams));
125
126        Ok(Self {
127            config: Arc::new(RwLock::new(config)),
128            oauth_state: Arc::new(oauth_state),
129            agent_state,
130            ai_service,
131            stream_semaphore,
132            active_tasks: ActiveTasks::default(),
133            cors,
134            port,
135        })
136    }
137
138    pub fn create_router(&self) -> Router {
139        let state = Arc::new(AgentHandlerState {
140            config: Arc::clone(&self.config),
141            oauth_state: Arc::clone(&self.oauth_state),
142            agent_state: Arc::clone(&self.agent_state),
143            ai_service: Arc::clone(&self.ai_service),
144            stream_semaphore: Arc::clone(&self.stream_semaphore),
145            active_tasks: self.active_tasks.clone(),
146        });
147
148        let post_router = Router::new()
149            .route("/", post(handle_agent_request))
150            .layer(DefaultBodyLimit::max(A2A_MAX_REQUEST_BODY_BYTES))
151            .with_state(Arc::clone(&state))
152            .layer(middleware::from_fn_with_state(
153                Arc::clone(&state),
154                agent_oauth_middleware_wrapper,
155            ));
156
157        let get_router = Router::new()
158            .route(ApiPaths::WELLKNOWN_AGENT_CARD, get(handle_agent_card))
159            .route(ApiPaths::A2A_CARD, get(handle_agent_card))
160            .with_state(state);
161
162        Router::new()
163            .merge(post_router)
164            .merge(get_router)
165            .layer(self.cors.clone())
166    }
167
168    // Why: the listener stops accepting on `shutdown`, then the server waits
169    // for every stream worker it spawned so an in-flight task is persisted
170    // rather than torn down mid-write.
171    pub async fn run<F>(self, shutdown: F) -> Result<(), crate::error::AgentError>
172    where
173        F: Future<Output = ()> + Send + 'static,
174    {
175        let app = self.create_router();
176        let addr = format!("0.0.0.0:{}", self.port);
177        let listener = tokio::net::TcpListener::bind(&addr).await?;
178        tracing::info!(
179            addr = %addr,
180            max_concurrent_streams = self.stream_semaphore.available_permits(),
181            "A2A server listening"
182        );
183
184        axum::serve(listener, app)
185            .with_graceful_shutdown(shutdown)
186            .await
187            .map_err(|e| crate::error::AgentError::Server(e.to_string()))?;
188
189        self.active_tasks.tracker().close();
190        self.active_tasks.tracker().wait().await;
191        Ok(())
192    }
193}