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