@@ -4,10 +4,10 @@ use axum::{
44 routing:: { delete, get, patch, post, put} ,
55} ;
66
7- use http:: { HeaderMap , HeaderName , StatusCode , Uri , header} ;
7+ use anyhow:: { Context , bail} ;
8+ use http:: { HeaderMap , HeaderName , HeaderValue , StatusCode , Uri , header} ;
89use rust_embed:: RustEmbed ;
9- use tower_http:: cors;
10- use tower_http:: cors:: CorsLayer ;
10+ use tower_http:: cors:: { AllowHeaders , AllowMethods , AllowOrigin , Any , CorsLayer } ;
1111use utoipa:: OpenApi ;
1212use utoipa_swagger_ui:: SwaggerUi ;
1313
@@ -31,7 +31,7 @@ use crate::pipelines::{
3131} ;
3232use crate :: rest_utils:: not_found;
3333use crate :: udfs:: { create_udf, delete_udf, get_udfs, validate_udf} ;
34- use arroyo_rpc:: config:: config;
34+ use arroyo_rpc:: config:: { CorsConfig , CorsOriginPolicy , config} ;
3535use cornucopia_async:: DatabaseSource ;
3636
3737static BASENAME_HEADER : HeaderName = HeaderName :: from_static ( "x-arroyo-basename" ) ;
@@ -156,12 +156,54 @@ async fn index_html(headers: HeaderMap) -> Response {
156156 }
157157}
158158
159- pub fn create_rest_app ( database : DatabaseSource ) -> Router {
160- // TODO: enable in development only!!!
159+ fn cors_layer ( config : & CorsConfig ) -> anyhow:: Result < CorsLayer > {
161160 let cors = CorsLayer :: new ( )
162- . allow_methods ( cors:: Any )
163- . allow_headers ( cors:: Any )
164- . allow_origin ( cors:: Any ) ;
161+ . allow_methods ( AllowMethods :: mirror_request ( ) )
162+ . allow_headers ( AllowHeaders :: mirror_request ( ) ) ;
163+
164+ match config. origin_policy {
165+ CorsOriginPolicy :: Any => {
166+ if config. allow_credentials {
167+ bail ! ( "CORS credentials cannot be enabled when allowing any origin" ) ;
168+ }
169+ if !config. allowed_origins . is_empty ( ) {
170+ bail ! ( "CORS allowed origins must be empty when allowing any origin" ) ;
171+ }
172+
173+ Ok ( cors. allow_credentials ( false ) . allow_origin ( Any ) )
174+ }
175+ CorsOriginPolicy :: AllowList => {
176+ if config. allowed_origins . iter ( ) . any ( |origin| origin == "*" ) {
177+ bail ! ( "CORS allow-list cannot contain the wildcard origin '*'" ) ;
178+ }
179+
180+ let origins = config
181+ . allowed_origins
182+ . iter ( )
183+ . map ( |origin| {
184+ let url = url:: Url :: parse ( origin)
185+ . with_context ( || format ! ( "invalid CORS allowed origin {origin:?}" ) ) ?;
186+ if url. origin ( ) . ascii_serialization ( ) != * origin {
187+ bail ! (
188+ "CORS allowed origin {origin:?} must contain only a canonical scheme, host, and optional port"
189+ ) ;
190+ }
191+
192+ origin
193+ . parse :: < HeaderValue > ( )
194+ . with_context ( || format ! ( "invalid CORS allowed origin {origin:?}" ) )
195+ } )
196+ . collect :: < anyhow:: Result < Vec < _ > > > ( ) ?;
197+
198+ Ok ( cors
199+ . allow_credentials ( config. allow_credentials )
200+ . allow_origin ( AllowOrigin :: list ( origins) ) )
201+ }
202+ }
203+ }
204+
205+ pub fn create_rest_app ( database : DatabaseSource ) -> anyhow:: Result < Router > {
206+ let cors = cors_layer ( & config ( ) . api . cors ) ?;
165207
166208 let api_routes = Router :: new ( )
167209 . route ( "/ping" , get ( ping) )
@@ -189,7 +231,7 @@ pub fn create_rest_app(database: DatabaseSource) -> Router {
189231 . merge ( pipeline_and_job_routes ( ) )
190232 . fallback ( api_fallback) ;
191233
192- Router :: new ( )
234+ Ok ( Router :: new ( )
193235 . merge (
194236 SwaggerUi :: new ( "/api/v1/swagger-ui" )
195237 . url ( "/api/v1/api-docs/openapi.json" , ApiDoc :: openapi ( ) ) ,
@@ -201,5 +243,5 @@ pub fn create_rest_app(database: DatabaseSource) -> Router {
201243 )
202244 . fallback ( static_handler)
203245 . with_state ( AppState { database } )
204- . layer ( cors)
246+ . layer ( cors) )
205247}
0 commit comments