systemprompt_api/services/middleware/
cors.rs1use 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}