diff --git a/Cargo.lock b/Cargo.lock index b59e53b371..6e7326895e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5596,6 +5596,7 @@ version = "1.3.0-alpha.1" dependencies = [ "api", "arc-swap", + "arrow-flight", "async-stream", "async-trait", "auth", diff --git a/src/frontend/Cargo.toml b/src/frontend/Cargo.toml index 29ca224a55..7ad24a466e 100644 --- a/src/frontend/Cargo.toml +++ b/src/frontend/Cargo.toml @@ -88,6 +88,7 @@ tonic.workspace = true tracing.workspace = true [dev-dependencies] +arrow-flight.workspace = true catalog = { workspace = true, features = ["testing"] } common-test-util.workspace = true datatypes.workspace = true diff --git a/src/frontend/src/server.rs b/src/frontend/src/server.rs index f0dc226ba8..e112e5ae2a 100644 --- a/src/frontend/src/server.rs +++ b/src/frontend/src/server.rs @@ -60,6 +60,7 @@ where http_server_builder: Option, plugins: Plugins, flight_handler: Option, + internal_flight_handler: Option, pub server_memory_limiter: ServerMemoryLimiter, } @@ -82,6 +83,7 @@ where http_server_builder: None, plugins, flight_handler: None, + internal_flight_handler: None, server_memory_limiter, } } @@ -209,6 +211,13 @@ where } } + pub fn with_internal_flight_handler(self, flight_handler: FlightCraftRef) -> Self { + Self { + internal_flight_handler: Some(flight_handler), + ..self + } + } + fn build_grpc_server( &mut self, grpc: &GrpcOptions, @@ -244,11 +253,16 @@ where grpc.flight_compression, ); - // Use custom flight handler if provided, otherwise use the default GreptimeRequestHandler - let flight_handler = self - .flight_handler - .clone() - .unwrap_or_else(|| Arc::new(greptime_request_handler.clone()) as FlightCraftRef); + let default_flight_handler = Arc::new(greptime_request_handler.clone()) as FlightCraftRef; + let flight_handler = if external { + self.flight_handler + .clone() + .unwrap_or(default_flight_handler) + } else { + self.internal_flight_handler + .clone() + .unwrap_or(default_flight_handler) + }; let grpc_server = builder .name(name) @@ -448,17 +462,50 @@ fn parse_addr(addr: &str) -> Result { #[cfg(test)] mod tests { + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; use std::time::Duration; use api::v1::HealthCheckRequest; use api::v1::health_check_client::HealthCheckClient; use api::v1::meta::Role; + use arrow_flight::{FlightData, PutResult, Ticket}; + use async_trait::async_trait; + use auth::{UserProviderRef, static_user_provider_from_option}; + use client::{Client, Database}; use meta_client::client::MetaClientBuilder; use servers::grpc::GRPC_SERVER; + use servers::grpc::flight::{FlightCraft, FlightCraftRef, TonicStream}; + use tonic::{Code, Request, Response, Status, Streaming}; use super::*; use crate::instance::builder::FrontendBuilder; + struct CountingFlightCraft { + inner: FlightCraftRef, + do_get_calls: AtomicUsize, + do_put_calls: AtomicUsize, + } + + #[async_trait] + impl FlightCraft for CountingFlightCraft { + async fn do_get( + &self, + request: Request, + ) -> std::result::Result>, Status> { + self.do_get_calls.fetch_add(1, Ordering::SeqCst); + self.inner.do_get(request).await + } + + async fn do_put( + &self, + request: Request>, + ) -> std::result::Result>, Status> { + self.do_put_calls.fetch_add(1, Ordering::SeqCst); + self.inner.do_put(request).await + } + } + #[test] fn test_effective_http_timeout_for_pending_rows() { let cases = [ @@ -549,6 +596,104 @@ mod tests { } } + #[tokio::test] + async fn test_database_sql_authentication_differs_between_public_and_internal_grpc() { + let options = FrontendOptions { + http: HttpOptions { + addr: "127.0.0.1:0".to_string(), + ..Default::default() + }, + grpc: GrpcOptions::default().with_bind_addr("127.0.0.1:0"), + internal_grpc: Some(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 plugins = Plugins::new(); + let provider = + static_user_provider_from_option("static_user_provider:cmd:greptime=greptime").unwrap(); + plugins.insert::(Arc::new(provider)); + let public_flight_handler = Arc::new(GreptimeRequestHandler::new( + instance.clone(), + plugins.get::(), + None, + options.grpc.flight_compression, + )) as FlightCraftRef; + let internal_flight_handler = Arc::new(CountingFlightCraft { + inner: Arc::new(GreptimeRequestHandler::new( + instance.clone(), + None, + None, + options.grpc.flight_compression, + )), + do_get_calls: AtomicUsize::new(0), + do_put_calls: AtomicUsize::new(0), + }); + let internal_flight_handler_ref = internal_flight_handler.clone() as FlightCraftRef; + let mut services = Services::new(options, instance, plugins) + .with_flight_handler(public_flight_handler) + .with_internal_flight_handler(internal_flight_handler_ref) + .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_database = Database::new( + "greptime", + "public", + Client::with_urls([public_addr.to_string()]), + ); + let internal_database = Database::new( + "greptime", + "public", + Client::with_urls([internal_addr.to_string()]), + ); + + let internal_result = internal_database.sql("SELECT 1").await; + let put_result = internal_database + .do_put(Box::pin(futures::stream::empty())) + .await; + let public_result = public_database.sql("SELECT 1").await; + + services.shutdown_all().await.unwrap(); + + assert!(internal_result.is_ok()); + assert!(put_result.is_ok()); + assert_eq!( + 1, + internal_flight_handler.do_get_calls.load(Ordering::SeqCst) + ); + assert_eq!( + 1, + internal_flight_handler.do_put_calls.load(Ordering::SeqCst) + ); + assert_eq!( + Some(Code::Unauthenticated), + public_result + .as_ref() + .err() + .and_then(|err| err.tonic_code()) + ); + } + #[tokio::test] async fn test_services_builder_health_check_is_reachable() { // Arrange