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