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