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