systemprompt_agent/services/a2a_server/
server.rs1use 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
26pub 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 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}