Skip to main content

agent_client_protocol_http/
server.rs

1use 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/// Configuration for an [`AcpHttpServer`].
17///
18/// Start with [`Self::default`] and use the fluent setters to customize it.
19/// Fields remain public for reading and mutation, but this type is non-exhaustive
20/// so new options can be added without breaking callers.
21///
22/// ```
23/// use agent_client_protocol_http::{CorsOptions, ServerOptions};
24///
25/// let options = ServerOptions::default()
26///     .with_path("/agent")
27///     .with_cors(CorsOptions::allow_origins(["https://example.com"])?)
28///     .with_health_endpoint(false);
29/// assert_eq!(options.path, "/agent");
30/// assert!(!options.health_endpoint);
31/// # Ok::<(), axum::http::header::InvalidHeaderValue>(())
32/// ```
33///
34/// Struct literals, including struct update syntax, are not supported outside
35/// this crate:
36///
37/// ```compile_fail
38/// use agent_client_protocol_http::ServerOptions;
39///
40/// let options = ServerOptions {
41///     path: "/agent".into(),
42///     ..ServerOptions::default()
43/// };
44/// ```
45#[derive(Debug, Clone)]
46#[non_exhaustive]
47pub struct ServerOptions {
48    /// ACP endpoint path. Defaults to `/acp`.
49    pub path: String,
50    /// Browser origin policy. Defaults to [`CorsOptions::Disabled`].
51    pub cors: CorsOptions,
52    /// Whether to expose `GET /health`. Defaults to `true`.
53    pub health_endpoint: bool,
54}
55
56impl ServerOptions {
57    /// Set the ACP endpoint path.
58    #[must_use]
59    pub fn with_path(mut self, path: impl Into<String>) -> Self {
60        self.path = path.into();
61        self
62    }
63
64    /// Set the browser origin policy.
65    ///
66    /// CORS does not authenticate requests. The built-in CORS layer does not
67    /// allow the `Authorization` request header or enable credentialed CORS.
68    /// Authentication and custom credential policies belong to the host router.
69    #[must_use]
70    pub fn with_cors(mut self, cors: CorsOptions) -> Self {
71        self.cors = cors;
72        self
73    }
74
75    /// Enable or disable the `GET /health` endpoint.
76    #[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/// Browser origin policy for HTTP CORS and WebSocket upgrades.
94///
95/// This is not authentication: requests without an `Origin` header are accepted
96/// by the origin check. Enforce authentication in the host router.
97///
98/// Prefer the constructors; matches outside this crate must include a wildcard
99/// arm to accommodate future policies:
100///
101/// ```
102/// use agent_client_protocol_http::CorsOptions;
103///
104/// let policy = CorsOptions::disabled();
105/// let disabled = match policy {
106///     CorsOptions::Disabled => true,
107///     _ => false,
108/// };
109/// assert!(disabled);
110/// ```
111///
112/// ```compile_fail
113/// use agent_client_protocol_http::CorsOptions;
114///
115/// let policy = CorsOptions::disabled();
116/// match policy {
117///     CorsOptions::Disabled => {},
118///     CorsOptions::AllowOrigins(_) => {},
119///     CorsOptions::AllowAnyOrigin => {},
120/// }
121/// ```
122#[derive(Debug, Clone, Default, PartialEq, Eq)]
123#[non_exhaustive]
124pub enum CorsOptions {
125    /// Disable cross-origin HTTP access and reject WebSocket browser origins.
126    #[default]
127    Disabled,
128    /// Allow only the listed browser origins.
129    AllowOrigins(Vec<HeaderValue>),
130    /// Allow all browser origins. Does not enable credentialed CORS.
131    AllowAnyOrigin,
132}
133
134impl CorsOptions {
135    /// Disable cross-origin browser access.
136    #[must_use]
137    pub fn disabled() -> Self {
138        Self::Disabled
139    }
140
141    /// Allow any browser origin, without enabling credentialed CORS.
142    #[must_use]
143    pub fn allow_any_origin() -> Self {
144        Self::AllowAnyOrigin
145    }
146
147    /// Allow the given browser origins.
148    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}