agent_client_protocol_http/
server.rs1use std::sync::Arc;
2
3use agent_client_protocol::{Client, ConnectTo};
4use axum::{
5 Router,
6 extract::WebSocketUpgrade,
7 extract::ws::rejection::WebSocketUpgradeRejection,
8 http::{HeaderName, HeaderValue, Method, StatusCode, header, header::InvalidHeaderValue},
9 response::{IntoResponse, Response},
10 routing::{delete, get, post},
11};
12use tower_http::cors::{AllowOrigin, CorsLayer};
13
14use crate::connection::ConnectionRegistry;
15
16#[derive(Debug, Clone)]
46#[non_exhaustive]
47pub struct ServerOptions {
48 pub path: String,
50 pub cors: CorsOptions,
52 pub health_endpoint: bool,
54}
55
56impl ServerOptions {
57 #[must_use]
59 pub fn with_path(mut self, path: impl Into<String>) -> Self {
60 self.path = path.into();
61 self
62 }
63
64 #[must_use]
70 pub fn with_cors(mut self, cors: CorsOptions) -> Self {
71 self.cors = cors;
72 self
73 }
74
75 #[must_use]
77 pub fn with_health_endpoint(mut self, enabled: bool) -> Self {
78 self.health_endpoint = enabled;
79 self
80 }
81}
82
83impl Default for ServerOptions {
84 fn default() -> Self {
85 Self {
86 path: "/acp".to_string(),
87 cors: CorsOptions::default(),
88 health_endpoint: true,
89 }
90 }
91}
92
93#[derive(Debug, Clone, Default, PartialEq, Eq)]
123#[non_exhaustive]
124pub enum CorsOptions {
125 #[default]
127 Disabled,
128 AllowOrigins(Vec<HeaderValue>),
130 AllowAnyOrigin,
132}
133
134impl CorsOptions {
135 #[must_use]
137 pub fn disabled() -> Self {
138 Self::Disabled
139 }
140
141 #[must_use]
143 pub fn allow_any_origin() -> Self {
144 Self::AllowAnyOrigin
145 }
146
147 pub fn allow_origins<I, S>(origins: I) -> Result<Self, InvalidHeaderValue>
149 where
150 I: IntoIterator<Item = S>,
151 S: AsRef<str>,
152 {
153 origins
154 .into_iter()
155 .map(|origin| HeaderValue::from_str(origin.as_ref()))
156 .collect::<Result<Vec<_>, _>>()
157 .map(Self::AllowOrigins)
158 }
159
160 fn allow_origin_layer(&self) -> Option<AllowOrigin> {
161 match self {
162 Self::Disabled => None,
163 Self::AllowOrigins(origins) => Some(AllowOrigin::list(origins.clone())),
164 Self::AllowAnyOrigin => Some(AllowOrigin::any()),
165 }
166 }
167
168 fn allows_origin(&self, origin: Option<&HeaderValue>) -> bool {
169 let Some(origin) = origin else {
170 return true;
171 };
172 match self {
173 Self::Disabled => false,
174 Self::AllowOrigins(origins) => origins.iter().any(|allowed| allowed == origin),
175 Self::AllowAnyOrigin => true,
176 }
177 }
178}
179
180#[derive(Clone)]
181struct ServerState {
182 registry: Arc<ConnectionRegistry>,
183 cors: CorsOptions,
184}
185
186pub struct AcpHttpServer {
187 registry: Arc<ConnectionRegistry>,
188 options: ServerOptions,
189}
190
191impl std::fmt::Debug for AcpHttpServer {
192 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
193 f.debug_struct("AcpHttpServer")
194 .field("options", &self.options)
195 .finish_non_exhaustive()
196 }
197}
198
199impl AcpHttpServer {
200 pub fn new<F, C>(factory: F) -> Self
201 where
202 F: Fn() -> C + Send + Sync + 'static,
203 C: ConnectTo<Client>,
204 {
205 Self {
206 registry: Arc::new(ConnectionRegistry::new(Arc::new(factory))),
207 options: ServerOptions::default(),
208 }
209 }
210
211 #[must_use]
212 pub fn with_options(mut self, options: ServerOptions) -> Self {
213 self.options = options;
214 self
215 }
216
217 pub fn into_router(self) -> Router {
218 let registry = self.registry.clone();
219 let path = self.options.path.clone();
220 let cors = self.options.cors.clone();
221 let state = ServerState {
222 registry: registry.clone(),
223 cors: cors.clone(),
224 };
225
226 let mut router = Router::new()
227 .route(
228 &path,
229 post(crate::http_server::handle_post).with_state(registry.clone()),
230 )
231 .route(&path, get(handle_get).with_state(state))
232 .route(
233 &path,
234 delete(crate::http_server::handle_delete).with_state(registry),
235 );
236
237 if self.options.health_endpoint {
238 router = router.route("/health", get(health));
239 }
240
241 if let Some(allow_origin) = cors.allow_origin_layer() {
242 router = router.layer(default_cors(allow_origin));
243 }
244
245 router
246 }
247}
248
249async fn health() -> &'static str {
250 "ok"
251}
252
253fn default_cors(allow_origin: AllowOrigin) -> CorsLayer {
254 CorsLayer::new()
255 .allow_origin(allow_origin)
256 .allow_methods([Method::GET, Method::POST, Method::DELETE, Method::OPTIONS])
257 .allow_headers([
258 header::CONTENT_TYPE,
259 header::ACCEPT,
260 HeaderName::from_static("acp-connection-id"),
261 HeaderName::from_static("acp-session-id"),
262 header::SEC_WEBSOCKET_VERSION,
263 header::SEC_WEBSOCKET_KEY,
264 header::CONNECTION,
265 header::UPGRADE,
266 ])
267 .expose_headers([
268 HeaderName::from_static("acp-connection-id"),
269 HeaderName::from_static("acp-session-id"),
270 ])
271}
272
273async fn handle_get(
274 ws_upgrade: Result<WebSocketUpgrade, WebSocketUpgradeRejection>,
275 axum::extract::State(state): axum::extract::State<ServerState>,
276 request: axum::http::Request<axum::body::Body>,
277) -> Response {
278 match ws_upgrade {
279 Ok(ws) => {
280 if !state
281 .cors
282 .allows_origin(request.headers().get(header::ORIGIN))
283 {
284 return (StatusCode::FORBIDDEN, "WebSocket origin not allowed").into_response();
285 }
286 crate::websocket_server::handle_ws_upgrade(state.registry, ws)
287 }
288 Err(_) => crate::http_server::handle_get(state.registry, request).await,
289 }
290}
291
292#[cfg(test)]
293mod tests {
294 use super::*;
295 use axum::body::Body;
296 use tower::{Layer as _, ServiceExt as _, service_fn};
297
298 #[test]
299 fn cors_is_disabled_by_default() {
300 assert_eq!(ServerOptions::default().cors, CorsOptions::Disabled);
301 }
302
303 #[test]
304 fn disabled_cors_rejects_browser_origin_for_websockets() {
305 let origin = HeaderValue::from_static("http://localhost:5173");
306
307 assert!(CorsOptions::disabled().allows_origin(None));
308 assert!(!CorsOptions::disabled().allows_origin(Some(&origin)));
309 }
310
311 #[test]
312 fn cors_allowlist_matches_configured_origins() {
313 let allowed = HeaderValue::from_static("http://localhost:5173");
314 let denied = HeaderValue::from_static("http://localhost:3000");
315 let cors = CorsOptions::allow_origins(["http://localhost:5173"]).unwrap();
316
317 assert!(cors.allows_origin(None));
318 assert!(cors.allows_origin(Some(&allowed)));
319 assert!(!cors.allows_origin(Some(&denied)));
320 }
321
322 #[test]
323 fn explicit_allow_any_origin_accepts_browser_origins() {
324 let origin = HeaderValue::from_static("https://example.com");
325
326 assert!(CorsOptions::allow_any_origin().allows_origin(Some(&origin)));
327 }
328
329 #[tokio::test]
330 async fn allow_any_origin_uses_wildcard_cors_header() {
331 let response = default_cors(
332 CorsOptions::allow_any_origin()
333 .allow_origin_layer()
334 .expect("CORS layer"),
335 )
336 .layer(service_fn(|_: axum::http::Request<Body>| async {
337 Ok::<_, std::convert::Infallible>(Response::new(Body::empty()))
338 }))
339 .oneshot(
340 axum::http::Request::builder()
341 .header(header::ORIGIN, "https://example.com")
342 .body(Body::empty())
343 .unwrap(),
344 )
345 .await
346 .unwrap();
347
348 assert_eq!(
349 response.headers().get(header::ACCESS_CONTROL_ALLOW_ORIGIN),
350 Some(&HeaderValue::from_static("*"))
351 );
352 assert!(response.headers().get(header::VARY).is_none());
353 }
354
355 #[tokio::test]
356 async fn allowlisted_origins_vary_by_origin() {
357 let response = default_cors(
358 CorsOptions::allow_origins(["https://example.com"])
359 .unwrap()
360 .allow_origin_layer()
361 .expect("CORS layer"),
362 )
363 .layer(service_fn(|_: axum::http::Request<Body>| async {
364 Ok::<_, std::convert::Infallible>(Response::new(Body::empty()))
365 }))
366 .oneshot(
367 axum::http::Request::builder()
368 .header(header::ORIGIN, "https://example.com")
369 .body(Body::empty())
370 .unwrap(),
371 )
372 .await
373 .unwrap();
374
375 assert_eq!(
376 response.headers().get(header::ACCESS_CONTROL_ALLOW_ORIGIN),
377 Some(&HeaderValue::from_static("https://example.com"))
378 );
379 assert_eq!(
380 response.headers().get(header::VARY),
381 Some(&HeaderValue::from_static("origin"))
382 );
383 }
384}