Skip to main content

systemprompt_api/services/middleware/
cors.rs

1//! CORS layer construction from profile config.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use axum::http::Method;
7use systemprompt_manifest::Config;
8use thiserror::Error;
9use tower_http::cors::{AllowOrigin, CorsLayer};
10
11#[derive(Debug, Error)]
12pub enum CorsError {
13    #[error("Invalid origin '{origin}' in cors_allowed_origins")]
14    InvalidOrigin {
15        origin: String,
16        #[source]
17        source: http::header::InvalidHeaderValue,
18    },
19    #[error("cors_allowed_origins must contain at least one valid origin")]
20    EmptyOrigins,
21}
22
23#[derive(Debug, Clone, Copy)]
24pub struct CorsMiddleware;
25
26impl CorsMiddleware {
27    pub fn build_layer(config: &Config) -> Result<CorsLayer, CorsError> {
28        let mut origins = Vec::new();
29        for origin in &config.cors_allowed_origins {
30            let trimmed = origin.trim();
31            if trimmed.is_empty() {
32                continue;
33            }
34            let header_value = trimmed.parse::<http::HeaderValue>().map_err(|source| {
35                CorsError::InvalidOrigin {
36                    origin: origin.clone(),
37                    source,
38                }
39            })?;
40            origins.push(header_value);
41        }
42
43        if origins.is_empty() {
44            return Err(CorsError::EmptyOrigins);
45        }
46
47        Ok(CorsLayer::new()
48            .allow_origin(AllowOrigin::list(origins))
49            .allow_credentials(true)
50            .allow_methods([
51                Method::GET,
52                Method::POST,
53                Method::PUT,
54                Method::DELETE,
55                Method::OPTIONS,
56            ])
57            .allow_headers([
58                http::header::AUTHORIZATION,
59                http::header::CONTENT_TYPE,
60                http::header::ACCEPT,
61                http::header::ORIGIN,
62                http::header::ACCESS_CONTROL_REQUEST_METHOD,
63                http::header::ACCESS_CONTROL_REQUEST_HEADERS,
64                http::HeaderName::from_static("mcp-protocol-version"),
65                http::HeaderName::from_static("x-context-id"),
66                http::HeaderName::from_static("x-gateway-conversation-id"),
67                http::HeaderName::from_static("x-provider-request-id"),
68                http::HeaderName::from_static("x-trace-id"),
69                http::HeaderName::from_static("x-call-source"),
70            ])
71            .expose_headers([
72                http::header::WWW_AUTHENTICATE,
73                http::HeaderName::from_static(systemprompt_gateway::service::RECOVERY_COUNT_HEADER),
74            ]))
75    }
76}