diff --git a/config/config.md b/config/config.md index 06d3b89e075..24351b2b6fd 100644 --- a/config/config.md +++ b/config/config.md @@ -43,6 +43,8 @@ | `grpc` | -- | -- | The gRPC server options. | | `grpc.bind_addr` | String | `127.0.0.1:4001` | The address to bind the gRPC server. | | `grpc.runtime_size` | Integer | `8` | The number of server worker threads. | +| `grpc.enable_cors` | Bool | `false` | Enable CORS for gRPC-Web clients in browsers. Disabled by default. | +| `grpc.cors_allowed_origins` | Array | Unset | Origins allowed by gRPC CORS. An empty list allows any origin. | | `grpc.max_connection_age` | String | Unset | The maximum connection age for gRPC connection.
The value can be a human-readable time string. For example: `10m` for ten minutes or `1h` for one hour.
Refer to https://grpc.io/docs/guides/keepalive/ for more details. | | `grpc.tls` | -- | -- | gRPC server TLS options, see `mysql.tls` section. | | `grpc.tls.mode` | String | `disable` | TLS mode. | @@ -303,6 +305,8 @@ | `grpc.server_addr` | String | `127.0.0.1:4001` | The address advertised to the metasrv, and used for connections from outside the host.
If left empty or unset, the server will automatically use the IP address of the first network interface
on the host, with the same port number as the one specified in `grpc.bind_addr`. | | `grpc.runtime_size` | Integer | `8` | The number of server worker threads. | | `grpc.flight_compression` | String | `arrow_ipc` | Compression mode for frontend side Arrow IPC service. Available options:
- `none`: disable all compression
- `transport`: only enable gRPC transport compression (zstd)
- `arrow_ipc`: only enable Arrow IPC compression (lz4)
- `all`: enable all compression.
Default to `none` | +| `grpc.enable_cors` | Bool | `false` | Enable CORS for gRPC-Web clients in browsers. Disabled by default. | +| `grpc.cors_allowed_origins` | Array | Unset | Origins allowed by gRPC CORS. An empty list allows any origin. | | `grpc.max_connection_age` | String | Unset | The maximum connection age for gRPC connection.
The value can be a human-readable time string. For example: `10m` for ten minutes or `1h` for one hour.
Refer to https://grpc.io/docs/guides/keepalive/ for more details. | | `grpc.tls` | -- | -- | gRPC server TLS options, see `mysql.tls` section. | | `grpc.tls.mode` | String | `disable` | TLS mode. | diff --git a/config/frontend.example.toml b/config/frontend.example.toml index b528da0304f..2f2c15cfc57 100644 --- a/config/frontend.example.toml +++ b/config/frontend.example.toml @@ -96,6 +96,11 @@ runtime_size = 8 ## - `all`: enable all compression. ## Default to `none` flight_compression = "arrow_ipc" +## Enable CORS for gRPC-Web clients in browsers. Disabled by default. +#+ enable_cors = false +## Origins allowed by gRPC CORS. An empty list allows any origin. +## @toml2docs:none-default +#+ cors_allowed_origins = ["https://example.com"] ## The maximum connection age for gRPC connection. ## The value can be a human-readable time string. For example: `10m` for ten minutes or `1h` for one hour. ## Refer to https://grpc.io/docs/guides/keepalive/ for more details. diff --git a/config/standalone.example.toml b/config/standalone.example.toml index c09c815aed0..4760e038f49 100644 --- a/config/standalone.example.toml +++ b/config/standalone.example.toml @@ -111,6 +111,11 @@ api_server_addr = "127.0.0.1:4006" bind_addr = "127.0.0.1:4001" ## The number of server worker threads. runtime_size = 8 +## Enable CORS for gRPC-Web clients in browsers. Disabled by default. +#+ enable_cors = false +## Origins allowed by gRPC CORS. An empty list allows any origin. +## @toml2docs:none-default +#+ cors_allowed_origins = ["https://example.com"] ## The maximum connection age for gRPC connection. ## The value can be a human-readable time string. For example: `10m` for ten minutes or `1h` for one hour. ## Refer to https://grpc.io/docs/guides/keepalive/ for more details. diff --git a/src/cmd/tests/load_config_test.rs b/src/cmd/tests/load_config_test.rs index cc095c5c92e..b56299f9916 100644 --- a/src/cmd/tests/load_config_test.rs +++ b/src/cmd/tests/load_config_test.rs @@ -464,6 +464,7 @@ fn test_load_standalone_example_config() { cors_allowed_origins: vec!["https://example.com".to_string()], ..Default::default() }, + grpc: GrpcOptions::default(), query: QueryOptions { memory_pool_size: MemoryLimit::Percentage(50), ..Default::default() diff --git a/src/frontend/src/frontend.rs b/src/frontend/src/frontend.rs index 71e5888b8fa..90e7abc4517 100644 --- a/src/frontend/src/frontend.rs +++ b/src/frontend/src/frontend.rs @@ -61,7 +61,8 @@ pub struct FrontendOptions { pub http: HttpOptions, pub grpc: GrpcOptions, /// The internal gRPC options for the frontend service. - /// it provide the same service as the public gRPC service, just only for internal use. + /// It serves the same services as the public one plus the internal handler. + /// CORS is always off on it. pub internal_grpc: Option, pub mysql: MysqlOptions, pub postgres: PostgresOptions, diff --git a/src/frontend/src/server.rs b/src/frontend/src/server.rs index a5568b24be5..cf278f1f1d0 100644 --- a/src/frontend/src/server.rs +++ b/src/frontend/src/server.rs @@ -215,11 +215,15 @@ where external: bool, request_memory_limiter: ServerMemoryLimiter, ) -> Result { - let builder = if let Some(builder) = self.grpc_server_builder.take() { + let mut builder = if let Some(builder) = self.grpc_server_builder.take() { builder } else { self.grpc_server_builder(grpc, request_memory_limiter)? }; + // Browsers never talk to the internal server. + if external && grpc.enable_cors { + builder = builder.with_cors(grpc.cors_allowed_origins.clone()); + } let user_provider = if external { self.plugins.get::() @@ -522,6 +526,7 @@ mod tests { use client::{Client, Database}; use common_grpc::channel_manager::ChannelManager; use meta_client::client::MetaClientBuilder; + use reqwest::header::{ACCESS_CONTROL_ALLOW_ORIGIN, ACCESS_CONTROL_REQUEST_METHOD, ORIGIN}; use servers::grpc::GRPC_SERVER; use servers::grpc::flight::{FlightCraft, FlightCraftRef, TonicStream}; use tonic::{Code, Request, Response, Status, Streaming}; @@ -1029,4 +1034,68 @@ mod tests { // Assert assert!(health_check.is_ok()); } + + #[tokio::test] + async fn test_internal_grpc_server_never_serves_cors() { + let options = FrontendOptions { + http: HttpOptions { + addr: "127.0.0.1:0".to_string(), + ..Default::default() + }, + grpc: GrpcOptions { + enable_cors: true, + ..GrpcOptions::default().with_bind_addr("127.0.0.1:0") + }, + internal_grpc: Some(GrpcOptions { + enable_cors: true, + ..GrpcOptions::default().with_bind_addr("127.0.0.1:0") + }), + mysql: crate::service_config::MysqlOptions { + enable: false, + ..Default::default() + }, + postgres: crate::service_config::PostgresOptions { + enable: false, + ..Default::default() + }, + ..Default::default() + }; + let meta_client = Arc::new( + MetaClientBuilder::new(0, Role::Frontend) + .enable_procedure() + .build(), + ); + let instance = Arc::new( + FrontendBuilder::new_test(&options, meta_client) + .try_build() + .await + .unwrap(), + ); + let mut services = Services::new(options, instance, Default::default()) + .build() + .unwrap(); + + services.start_all().await.unwrap(); + let public_addr = services.addr(GRPC_SERVER).unwrap(); + let internal_addr = services.addr("INTERNAL_GRPC_SERVER").unwrap(); + let public = send_cors_preflight(public_addr).await; + let internal = send_cors_preflight(internal_addr).await; + services.shutdown_all().await.unwrap(); + + assert!(public.headers().contains_key(ACCESS_CONTROL_ALLOW_ORIGIN)); + assert!(!internal.headers().contains_key(ACCESS_CONTROL_ALLOW_ORIGIN)); + } + + async fn send_cors_preflight(addr: std::net::SocketAddr) -> reqwest::Response { + reqwest::Client::new() + .request( + reqwest::Method::OPTIONS, + format!("http://{addr}/greptime.v1.HealthCheck/Check"), + ) + .header(ORIGIN, "https://example.com") + .header(ACCESS_CONTROL_REQUEST_METHOD, "POST") + .send() + .await + .unwrap() + } } diff --git a/src/servers/src/grpc.rs b/src/servers/src/grpc.rs index 1be1f8215b5..9e9335b422d 100644 --- a/src/servers/src/grpc.rs +++ b/src/servers/src/grpc.rs @@ -36,6 +36,7 @@ use common_grpc::channel_manager::{ }; use common_telemetry::{error, info, warn}; use futures::FutureExt; +use http::{HeaderName, Method}; use otel_arrow_rust::proto::opentelemetry::arrow::v1::arrow_metrics_service_server::ArrowMetricsServiceServer; use serde::{Deserialize, Serialize}; use snafu::{OptionExt, ResultExt, ensure}; @@ -48,9 +49,11 @@ use tonic::transport::ServerTlsConfig; use tonic::transport::server::TcpIncoming; use tonic::{Request, Response, Status}; use tonic_reflection::server::v1::{ServerReflection, ServerReflectionServer}; +use tower_http::cors::{AllowHeaders, CorsLayer}; use crate::error::{AlreadyStartedSnafu, InternalSnafu, Result, StartGrpcSnafu, TcpBindSnafu}; use crate::grpc::memory_limit::MemoryLimiterExtensionService; +use crate::http::cors_allow_origin; use crate::install_default_crypto_provider; use crate::metrics::MetricsMiddlewareLayer; use crate::otel_arrow::{HeaderInterceptor, OtelArrowServiceHandler}; @@ -87,6 +90,11 @@ pub struct GrpcOptions { /// The HTTP/2 keep-alive timeout. #[serde(with = "humantime_serde")] pub http2_keep_alive_timeout: Duration, + /// Whether to enable CORS, required by gRPC-Web clients in browsers. + /// Only the frontend's public gRPC server honors it. + pub enable_cors: bool, + /// Origins allowed by CORS. Empty allows any origin. + pub cors_allowed_origins: Vec, } impl GrpcOptions { @@ -155,6 +163,8 @@ impl Default for GrpcOptions { max_connection_age: None, http2_keep_alive_interval: Duration::from_secs(10), http2_keep_alive_timeout: Duration::from_secs(3), + enable_cors: false, + cors_allowed_origins: Vec::new(), } } } @@ -176,6 +186,8 @@ impl GrpcOptions { max_connection_age: None, http2_keep_alive_interval: Duration::from_secs(10), http2_keep_alive_timeout: Duration::from_secs(3), + enable_cors: false, + cors_allowed_origins: Vec::new(), } } @@ -237,6 +249,8 @@ pub struct GrpcServer { bind_addr: Option, name: Option, config: GrpcServerConfig, + /// `None` disables CORS. + cors_allowed_origins: Option>, } /// Grpc Server configuration @@ -265,6 +279,24 @@ impl Default for GrpcServerConfig { } impl GrpcServer { + fn cors_layer(&self) -> Result> { + let Some(allowed_origins) = &self.cors_allowed_origins else { + return Ok(None); + }; + let layer = CorsLayer::new() + .allow_methods([Method::POST]) + .allow_origin(cors_allow_origin(allowed_origins)?) + .allow_headers(AllowHeaders::any()) + // Trailers-only responses (most errors) carry `grpc-status` in HTTP + // headers, which browsers hide unless exposed. + .expose_headers([ + HeaderName::from_static("grpc-status"), + HeaderName::from_static("grpc-message"), + HeaderName::from_static("grpc-status-details-bin"), + ]); + Ok(Some(layer)) + } + pub fn create_healthcheck_service(&self) -> HealthCheckServer { HealthCheckServer::new(HealthCheckHandler) } @@ -359,13 +391,18 @@ impl Server for GrpcServer { (incoming, addr) }; - let metrics_layer = tower::ServiceBuilder::new() + let cors_layer = self.cors_layer()?; + + // `Server::layer` prepends, so the first layer added is outermost. CORS must + // sit outside `GrpcWebLayer`, which rejects the `OPTIONS` preflight with a 400. + let middleware_layer = tower::ServiceBuilder::new() .layer(MetricsMiddlewareLayer) + .option_layer(cors_layer) .into_inner(); let mut builder = tonic::transport::Server::builder() .accept_http1(true) - .layer(metrics_layer) + .layer(middleware_layer) .layer(tonic_web::GrpcWebLayer::new()); if let Some(tls_config) = self.tls_config.clone() { @@ -425,9 +462,23 @@ impl Server for GrpcServer { #[cfg(test)] mod tests { - use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; + use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; - use super::{DEFAULT_GRPC_ADDR_PORT, format_server_addr, port_from_bind_addr}; + use common_runtime::Runtime; + use common_runtime::runtime::{BuilderBuild, RuntimeTrait}; + use http::header::{ + ACCESS_CONTROL_ALLOW_HEADERS, ACCESS_CONTROL_ALLOW_METHODS, ACCESS_CONTROL_ALLOW_ORIGIN, + ACCESS_CONTROL_EXPOSE_HEADERS, ACCESS_CONTROL_REQUEST_HEADERS, + ACCESS_CONTROL_REQUEST_METHOD, CONTENT_TYPE, ORIGIN, + }; + use http::{HeaderName, Method, StatusCode}; + + use super::{ + DEFAULT_GRPC_ADDR_PORT, GrpcServer, GrpcServerConfig, format_server_addr, + port_from_bind_addr, + }; + use crate::grpc::builder::GrpcServerBuilder; + use crate::server::Server; #[test] fn test_port_from_bind_addr() { @@ -454,4 +505,118 @@ mod tests { format_server_addr(IpAddr::V6(Ipv6Addr::LOCALHOST), 3002) ); } + + /// The returned server owns the shutdown sender and must outlive the requests. + async fn start_test_grpc_server( + cors_allowed_origins: Option>, + ) -> (GrpcServer, SocketAddr) { + let runtime = Runtime::builder().build().unwrap(); + let mut builder = GrpcServerBuilder::new(GrpcServerConfig::default(), runtime); + if let Some(origins) = cors_allowed_origins { + builder = builder.with_cors(origins); + } + let mut server = builder.build(); + server + .start(SocketAddr::from((Ipv4Addr::LOCALHOST, 0))) + .await + .unwrap(); + let addr = server.bind_addr().unwrap(); + (server, addr) + } + + const TEST_PATH: &str = "greptime.v1.HealthCheck/Check"; + + async fn send_preflight(addr: SocketAddr, origin: &str) -> reqwest::Response { + reqwest::Client::new() + .request(Method::OPTIONS, format!("http://{addr}/{TEST_PATH}")) + .header(ORIGIN, origin) + .header(ACCESS_CONTROL_REQUEST_METHOD, "POST") + .header(ACCESS_CONTROL_REQUEST_HEADERS, "content-type,x-grpc-web") + .send() + .await + .unwrap() + } + + /// The path is not served, so the response is a trailers-only `Unimplemented`. + async fn send_grpc_web_call(addr: SocketAddr, origin: &str) -> reqwest::Response { + reqwest::Client::new() + .post(format!("http://{addr}/{TEST_PATH}")) + .header(ORIGIN, origin) + .header(CONTENT_TYPE, "application/grpc-web+proto") + .header("x-grpc-web", "1") + // Empty gRPC-Web frame: 1-byte compression flag + 4-byte length. + .body(vec![0u8; 5]) + .send() + .await + .unwrap() + } + + fn header(response: &reqwest::Response, name: HeaderName) -> Option<&str> { + response.headers().get(name).map(|v| v.to_str().unwrap()) + } + + #[tokio::test] + async fn test_grpc_web_cors_preflight() { + let (_server, addr) = start_test_grpc_server(Some(Vec::new())).await; + + let response = send_preflight(addr, "https://example.com").await; + + assert_eq!(response.status(), StatusCode::OK); + assert_eq!(header(&response, ACCESS_CONTROL_ALLOW_ORIGIN), Some("*")); + assert_eq!(header(&response, ACCESS_CONTROL_ALLOW_HEADERS), Some("*")); + // Lists the method of the actual call, not `OPTIONS`. + assert_eq!( + header(&response, ACCESS_CONTROL_ALLOW_METHODS), + Some("POST") + ); + } + + #[tokio::test] + async fn test_grpc_web_cors_expose_headers() { + let (_server, addr) = start_test_grpc_server(Some(Vec::new())).await; + + let response = send_grpc_web_call(addr, "https://example.com").await; + + assert_eq!(response.status(), StatusCode::OK); + assert_eq!(header(&response, ACCESS_CONTROL_ALLOW_ORIGIN), Some("*")); + // Premise: the response is trailers-only. + assert_eq!( + header(&response, HeaderName::from_static("grpc-status")), + Some("12") + ); + let expose_headers = header(&response, ACCESS_CONTROL_EXPOSE_HEADERS).unwrap(); + for name in ["grpc-status", "grpc-message", "grpc-status-details-bin"] { + assert!(expose_headers.contains(name), "{expose_headers}"); + } + } + + #[tokio::test] + async fn test_grpc_web_cors_custom_origins() { + let (_server, addr) = + start_test_grpc_server(Some(vec!["https://example.com".to_string()])).await; + + let response = send_preflight(addr, "https://example.com").await; + assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + header(&response, ACCESS_CONTROL_ALLOW_ORIGIN), + Some("https://example.com") + ); + + // A disallowed origin still gets a 200, just without allow-origin; the + // browser does the blocking. + let response = send_preflight(addr, "https://notallowed.com").await; + assert_eq!(response.status(), StatusCode::OK); + assert_eq!(header(&response, ACCESS_CONTROL_ALLOW_ORIGIN), None); + } + + #[tokio::test] + async fn test_grpc_web_cors_disabled() { + let (_server, addr) = start_test_grpc_server(None).await; + + let response = send_preflight(addr, "https://example.com").await; + + // Without CORS, `GrpcWebLayer` rejects the preflight. + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!(header(&response, ACCESS_CONTROL_ALLOW_ORIGIN), None); + } } diff --git a/src/servers/src/grpc/builder.rs b/src/servers/src/grpc/builder.rs index 99513dcceb3..eaa7b628961 100644 --- a/src/servers/src/grpc/builder.rs +++ b/src/servers/src/grpc/builder.rs @@ -97,6 +97,7 @@ pub struct GrpcServerBuilder { >, >, memory_limiter: ServerMemoryLimiter, + cors_allowed_origins: Option>, } impl GrpcServerBuilder { @@ -112,9 +113,16 @@ impl GrpcServerBuilder { tls_config: None, otel_arrow_service: None, memory_limiter, + cors_allowed_origins: None, } } + /// Enables CORS for gRPC-Web clients in browsers. An empty list allows any origin. + pub fn with_cors(mut self, allowed_origins: Vec) -> Self { + self.cors_allowed_origins = Some(allowed_origins); + self + } + /// Set a global memory limiter for all server protocols. pub fn with_memory_limiter(mut self, limiter: ServerMemoryLimiter) -> Self { self.memory_limiter = limiter; @@ -257,6 +265,7 @@ impl GrpcServerBuilder { bind_addr: None, name: self.name, config: self.config, + cors_allowed_origins: self.cors_allowed_origins, } } } diff --git a/src/servers/src/http.rs b/src/servers/src/http.rs index a33e52a1d40..e7cb715fa40 100644 --- a/src/servers/src/http.rs +++ b/src/servers/src/http.rs @@ -1061,20 +1061,7 @@ impl HttpServer { Method::DELETE, Method::HEAD, ]) - .allow_origin(if self.options.cors_allowed_origins.is_empty() { - AllowOrigin::from(Any) - } else { - AllowOrigin::from( - self.options - .cors_allowed_origins - .iter() - .map(|s| { - HeaderValue::from_str(s.as_str()) - .context(InvalidHeaderValueSnafu) - }) - .collect::>>()?, - ) - }) + .allow_origin(cors_allow_origin(&self.options.cors_allowed_origins)?) .allow_headers(Any), ) } else { @@ -1612,6 +1599,18 @@ impl HttpServer { pub const HTTP_SERVER: &str = "HTTP_SERVER"; pub const HTTP_API_SERVER: &str = "HTTP_API_SERVER"; +/// An empty list allows any origin. +pub(crate) fn cors_allow_origin(origins: &[String]) -> Result { + if origins.is_empty() { + return Ok(AllowOrigin::any()); + } + let origins = origins + .iter() + .map(|origin| HeaderValue::from_str(origin).context(InvalidHeaderValueSnafu)) + .collect::>>()?; + Ok(AllowOrigin::list(origins)) +} + #[async_trait] impl Server for HttpServer { async fn shutdown(&self) -> Result<()> { diff --git a/tests-integration/tests/http.rs b/tests-integration/tests/http.rs index 9b6ef9d8ae2..b1cfeac1ba1 100644 --- a/tests-integration/tests/http.rs +++ b/tests-integration/tests/http.rs @@ -2835,6 +2835,8 @@ flight_compression = "arrow_ipc" runtime_size = 8 http2_keep_alive_interval = "10s" http2_keep_alive_timeout = "3s" +enable_cors = false +cors_allowed_origins = [] [grpc.tls] mode = "disable"