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
7use axum::routing::{get, post};
8use axum::{Router, middleware};
9use std::pin::Pin;
10use std::sync::Arc;
11use systemprompt_database::DbPool;
12use systemprompt_models::modules::ApiPaths;
13use systemprompt_models::{AgentConfig, AiProvider};
14use tokio::sync::{RwLock, Semaphore};
15use tower_http::cors::CorsLayer;
16use tower_http::services::ServeDir;
17
18use super::auth::{AgentOAuthConfig, AgentOAuthState, agent_oauth_middleware_wrapper};
19use super::handlers::{AgentHandlerState, handle_agent_card, handle_agent_request};
20use crate::state::AgentState;
21
22pub struct Server {
23    db_pool: DbPool,
24    config: Arc<RwLock<AgentConfig>>,
25    oauth_state: Arc<AgentOAuthState>,
26    agent_state: Arc<AgentState>,
27    ai_service: Arc<dyn AiProvider>,
28    stream_semaphore: Arc<Semaphore>,
29    port: u16,
30}
31
32impl std::fmt::Debug for Server {
33    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
34        f.debug_struct("Server")
35            .field("db_pool", &"<DbPool>")
36            .field("config", &"Arc<RwLock<AgentConfig>>")
37            .field("oauth_state", &"Arc<AgentOAuthState>")
38            .field("agent_state", &"Arc<AgentState>")
39            .field("ai_service", &"<Arc<dyn AiProvider>>")
40            .field(
41                "stream_semaphore",
42                &self.stream_semaphore.available_permits(),
43            )
44            .field("port", &self.port)
45            .finish()
46    }
47}
48
49impl Server {
50    pub async fn new(
51        db_pool: DbPool,
52        agent_state: Arc<AgentState>,
53        ai_service: Arc<dyn AiProvider>,
54        agent_name: Option<String>,
55        port: u16,
56    ) -> Result<Self, crate::error::AgentError> {
57        use crate::services::registry::AgentRegistry;
58
59        let mut config = if let Some(name) = agent_name {
60            let registry = AgentRegistry::new()
61                .map_err(|e| crate::error::AgentError::Server(e.to_string()))?;
62            registry
63                .get_agent(&name)
64                .await
65                .map_err(|e| crate::error::AgentError::Server(e.to_string()))?
66        } else {
67            return Err(crate::error::AgentError::Validation(
68                "Agent name is required".to_owned(),
69            ));
70        };
71
72        config.extract_oauth_scopes_from_card();
73
74        let oauth_config = AgentOAuthConfig::default();
75        let global_config = systemprompt_models::Config::get()
76            .map_err(|e| crate::error::AgentError::Config(e.to_string()))?;
77        let mut oauth_state = AgentOAuthState::new(
78            Arc::clone(&db_pool),
79            oauth_config,
80            global_config.jwt_issuer.clone(),
81            global_config.jwt_audiences.clone(),
82        );
83
84        oauth_state = oauth_state.with_jwt_provider(Arc::clone(agent_state.jwt_provider()));
85
86        Ok(Self {
87            db_pool,
88            config: Arc::new(RwLock::new(config)),
89            oauth_state: Arc::new(oauth_state),
90            agent_state,
91            ai_service,
92            stream_semaphore: Arc::new(Semaphore::new(global_config.max_concurrent_streams)),
93            port,
94        })
95    }
96
97    pub fn create_router(&self) -> Router {
98        let state = Arc::new(AgentHandlerState {
99            db_pool: Arc::clone(&self.db_pool),
100            config: Arc::clone(&self.config),
101            oauth_state: Arc::clone(&self.oauth_state),
102            agent_state: Arc::clone(&self.agent_state),
103            ai_service: Arc::clone(&self.ai_service),
104            stream_semaphore: Arc::clone(&self.stream_semaphore),
105        });
106
107        let post_router = Router::new()
108            .route("/", post(handle_agent_request))
109            .with_state(Arc::clone(&state))
110            .layer(middleware::from_fn_with_state(
111                Arc::clone(&state),
112                agent_oauth_middleware_wrapper,
113            ));
114
115        let get_router = Router::new()
116            .route(ApiPaths::WELLKNOWN_AGENT_CARD, get(handle_agent_card))
117            .route(ApiPaths::A2A_CARD, get(handle_agent_card))
118            .with_state(state);
119
120        let api_router = Router::new().merge(post_router).merge(get_router);
121
122        let web_dist_path = std::path::Path::new("web/dist");
123        let router = if web_dist_path.exists() {
124            api_router.fallback_service(ServeDir::new(web_dist_path))
125        } else {
126            api_router
127        };
128
129        router.layer(CorsLayer::permissive())
130    }
131
132    pub async fn run(self) -> Result<(), crate::error::AgentError> {
133        Self::log_server_configuration();
134        self.start_server(None).await
135    }
136
137    const fn log_server_configuration() {}
138
139    async fn start_server(
140        self,
141        shutdown_signal: Option<Pin<Box<dyn Future<Output = ()> + Send>>>,
142    ) -> Result<(), crate::error::AgentError> {
143        let app = self.create_router();
144        let addr = format!("0.0.0.0:{}", self.port);
145        let listener = tokio::net::TcpListener::bind(&addr).await?;
146
147        match shutdown_signal {
148            Some(signal) => axum::serve(listener, app)
149                .with_graceful_shutdown(signal)
150                .await
151                .map_err(|e| crate::error::AgentError::Server(e.to_string())),
152            None => axum::serve(listener, app)
153                .await
154                .map_err(|e| crate::error::AgentError::Server(e.to_string())),
155        }
156    }
157}