diff --git a/Cargo.lock b/Cargo.lock index b5e21ec266..f0256e9392 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4727,6 +4727,7 @@ dependencies = [ "toml", "tonic", "tonic-prost-build", + "tonic-reflection", "tower", "tower-http", "tracing", @@ -7775,6 +7776,20 @@ dependencies = [ "tonic-build", ] +[[package]] +name = "tonic-reflection" +version = "0.14.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "acccd136a4bf19810a1fde9c74edc6129b42a66b44d0c1c8aaa67aeb49a146a7" +dependencies = [ + "prost", + "prost-types", + "tokio", + "tokio-stream", + "tonic", + "tonic-prost", +] + [[package]] name = "tonic-types" version = "0.14.6" diff --git a/Cargo.toml b/Cargo.toml index 7cab66d291..57415625f4 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -19,6 +19,7 @@ tokio = { version = "1.43", features = ["full"] } # gRPC/Protobuf tonic = "0.14" tonic-types = "0.14" +tonic-reflection = "0.14" tonic-prost = "0.14" tonic-prost-build = "0.14" prost = "0.14" diff --git a/README.md b/README.md index ed5cc57ecc..60133557a8 100644 --- a/README.md +++ b/README.md @@ -49,6 +49,23 @@ The installer sets up the CLI and a local gateway. The default sandbox image is - [Tutorials](https://docs.nvidia.com/openshell/latest/tutorials/first-network-policy): step-by-step policy and agent walkthroughs. - [Prerelease and development builds](https://docs.nvidia.com/openshell/latest/about/installation#prerelease-and-development-builds): try an upcoming release or the latest commit on `main`. +### Test the gRPC API with grpcurl + +The gateway serves the gRPC reflection v1 protocol. After starting a local +plaintext gateway, use `grpcurl` without checking out or supplying the proto +files: + +```shell +grpcurl -plaintext localhost:18080 list +grpcurl -plaintext localhost:18080 describe openshell.v1.OpenShell +grpcurl -plaintext -d '{}' localhost:18080 openshell.v1.OpenShell/Health +``` + +The service list contains the public `openshell.v1.OpenShell` API. Reflection +does not advertise the gateway's internal compute-driver, credential-driver, +interceptor, or middleware services. For a TLS gateway, omit `-plaintext` and +supply the CA and client certificate options required by the deployment. + ## Agent Skills Install the public OpenShell skills for your coding agent: diff --git a/crates/openshell-server/Cargo.toml b/crates/openshell-server/Cargo.toml index 54c44a7bed..7528a636b9 100644 --- a/crates/openshell-server/Cargo.toml +++ b/crates/openshell-server/Cargo.toml @@ -39,6 +39,7 @@ libc = "0.2" # gRPC tonic = { workspace = true, features = ["channel", "tls-native-roots"] } +tonic-reflection = { workspace = true } prost = { workspace = true } prost-reflect = { workspace = true } prost-types = { workspace = true } diff --git a/crates/openshell-server/src/auth/oidc.rs b/crates/openshell-server/src/auth/oidc.rs index 002385de95..962fab6e7a 100644 --- a/crates/openshell-server/src/auth/oidc.rs +++ b/crates/openshell-server/src/auth/oidc.rs @@ -32,7 +32,8 @@ use tracing::{debug, error, info, warn}; /// These are structural bypasses for gRPC infrastructure that doesn't map to a /// single RPC method. Per-method bypasses (e.g. `Health`) are declared at the /// handler with `auth_mode: "unauthenticated"` in the proto annotation. -const UNAUTHENTICATED_PREFIXES: &[&str] = &["/grpc.reflection.", "/grpc.health."]; +const UNAUTHENTICATED_PREFIXES: &[&str] = + &[crate::multiplex::REFLECTION_PATH_PREFIX, "/grpc.health."]; /// Returns `true` if the method needs no authentication at all. pub fn is_unauthenticated_method(path: &str) -> bool { @@ -1209,10 +1210,13 @@ mod tests { #[test] fn reflection_is_unauthenticated() { assert!(is_unauthenticated_method( + "/grpc.reflection.v1.ServerReflection/ServerReflectionInfo" + )); + assert!(!is_unauthenticated_method( "/grpc.reflection.v1alpha.ServerReflection/ServerReflectionInfo" )); - assert!(is_unauthenticated_method( - "/grpc.reflection.v1.ServerReflection/ServerReflectionInfo" + assert!(!is_unauthenticated_method( + "/grpc.reflection.v2.ServerReflection/ServerReflectionInfo" )); } diff --git a/crates/openshell-server/src/lib.rs b/crates/openshell-server/src/lib.rs index d6d35f6ba3..b8071a948e 100644 --- a/crates/openshell-server/src/lib.rs +++ b/crates/openshell-server/src/lib.rs @@ -35,6 +35,7 @@ pub(crate) mod policy_store; mod provider_profile_sources; mod provider_refresh; mod readiness; +mod reflection; mod sandbox_index; mod sandbox_watch; mod service_routing; @@ -340,6 +341,9 @@ pub struct ServerState { /// Gateway-wide gRPC request rate limiter shared by every multiplex path. pub(crate) grpc_rate_limiter: Option, + /// Immutable public reflection index and its per-query rate limiter. + pub(crate) reflection_service: reflection::GatewayReflectionServer, + /// Per-sandbox bound on extension credential minting, which resolves the /// caller's effective policy on every request. pub(crate) extension_mint_limiter: auth::extension_mint_limit::ExtensionMintLimiter, @@ -416,6 +420,8 @@ impl ServerState { let replica_id = compute::lease::replica_id(); let peer_endpoint = derive_peer_endpoint(&config); let grpc_rate_limiter = multiplex::GrpcRateLimiter::from_config(&config); + let reflection_service = reflection::build_gateway_reflection_service(&config) + .expect("compiled public gateway descriptors must be valid"); let admin_role = config .oidc .as_ref() @@ -446,6 +452,7 @@ impl ServerState { compute_driver_authenticator: None, peer_authenticator: None, grpc_rate_limiter, + reflection_service, gateway_interceptors: None, provider_profile_sources: provider_profile_sources::ProviderProfileSources::with_default_sources(), @@ -2232,6 +2239,68 @@ mod tests { (listen_addr, shutdown_tx, handle, tls_dir) } + async fn start_plaintext_gateway_listener() + -> (SocketAddr, watch::Sender, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("failed to bind test listener"); + let listen_addr = listener.local_addr().expect("failed to read local addr"); + let state = test_state(listen_addr, false).await; + let service = MultiplexService::new(state); + let (shutdown_tx, shutdown_rx) = watch::channel(false); + let handle = tokio::spawn(serve_gateway_listener( + BoundGatewayListener { + listener, + address: listen_addr, + }, + service, + None, + false, + shutdown_rx, + )); + (listen_addr, shutdown_tx, handle) + } + + #[tokio::test] + async fn production_gateway_listener_serves_reflection() { + use tonic_reflection::pb::v1::{ + ServerReflectionRequest, server_reflection_client::ServerReflectionClient, + server_reflection_request::MessageRequest, server_reflection_response::MessageResponse, + }; + + let (addr, shutdown, handle) = start_plaintext_gateway_listener().await; + let channel = tonic::transport::Channel::from_shared(format!("http://{addr}")) + .unwrap() + .connect() + .await + .unwrap(); + let mut client = ServerReflectionClient::new(channel); + let request = ServerReflectionRequest { + host: String::new(), + message_request: Some(MessageRequest::ListServices(String::new())), + }; + let mut responses = client + .server_reflection_info(tokio_stream::iter([request])) + .await + .unwrap() + .into_inner(); + let response = responses.message().await.unwrap().unwrap(); + let Some(MessageResponse::ListServicesResponse(services)) = response.message_response + else { + panic!("expected a reflection list-services response"); + }; + assert_eq!( + services + .service + .iter() + .map(|service| service.name.as_str()) + .collect::>(), + ["openshell.v1.OpenShell"] + ); + + stop_listener(shutdown, handle).await; + } + async fn send_plain_http(addr: SocketAddr, request: String) -> String { let connect_addr: SocketAddr = format!("127.0.0.1:{}", addr.port()) .parse() diff --git a/crates/openshell-server/src/multiplex.rs b/crates/openshell-server/src/multiplex.rs index 8988542b1a..53de58b398 100644 --- a/crates/openshell-server/src/multiplex.rs +++ b/crates/openshell-server/src/multiplex.rs @@ -199,7 +199,6 @@ macro_rules! request_id_middleware { /// the largest payload and well within this cap under normal use. const MAX_GRPC_DECODE_SIZE: usize = 1_048_576; const MAX_INTERCEPTED_GRPC_BODY_SIZE: usize = MAX_GRPC_DECODE_SIZE + 5; - /// Concurrent HTTP/2 streams allowed per connection. Sits above the /// per-replica pending relay budget so pooled peer connections are bounded by /// the relay caps rather than by the transport. @@ -243,6 +242,7 @@ impl MultiplexService { self.state.gateway_interceptors.clone(), Some(self.state.clone()), ); + let reflection = self.state.reflection_service.clone(); let authz_policy = self.state.config.oidc.as_ref().map(|oidc| AuthzPolicy { admin_role: oidc.admin_role.clone(), user_role: oidc.user_role.clone(), @@ -250,7 +250,7 @@ impl MultiplexService { }); let authenticator_chain = build_authenticator_chain(&self.state); let grpc_service = AuthGrpcRouter::with_peer_identity( - openshell, + GrpcRouter::new(openshell, reflection), authenticator_chain, authz_policy, self.state @@ -736,7 +736,7 @@ impl GrpcRateLimiter { }) } - fn allow(&self) -> bool { + pub(crate) fn allow(&self) -> bool { let now = Instant::now(); let mut state = self .state @@ -857,6 +857,56 @@ where } } +/// Combined gRPC service that routes between `OpenShell` and reflection. +#[derive(Clone)] +pub struct GrpcRouter { + openshell: N, + reflection: R, +} + +impl GrpcRouter { + fn new(openshell: N, reflection: R) -> Self { + Self { + openshell, + reflection, + } + } +} + +pub const REFLECTION_PATH_PREFIX: &str = "/grpc.reflection.v1."; + +impl tower::Service> for GrpcRouter +where + N: tower::Service> + Clone + Send + 'static, + N::Response: Send, + N::Future: Send, + N::Error: Send, + R: tower::Service, Response = N::Response, Error = N::Error> + + Clone + + Send + + 'static, + R::Future: Send, + B: Send + 'static, +{ + type Response = N::Response; + type Error = N::Error; + type Future = Pin> + Send>>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, req: Request) -> Self::Future { + if req.uri().path().starts_with(REFLECTION_PATH_PREFIX) { + let mut svc = self.reflection.clone(); + Box::pin(async move { svc.ready().await?.call(req).await }) + } else { + let mut svc = self.openshell.clone(); + Box::pin(async move { svc.ready().await?.call(req).await }) + } + } +} + /// Assemble the authenticator chain for the gateway. /// /// Chain order (first-match-wins): @@ -2493,6 +2543,310 @@ mod tests { assert_eq!(grpc_method_from_path(""), ""); } + #[tokio::test] + async fn grpc_router_dispatches_gateway_and_reflection_paths() { + #[derive(Clone)] + struct RouteRecorder { + name: &'static str, + calls: Arc>>, + } + + impl Service> for RouteRecorder { + type Response = Response; + type Error = Infallible; + type Future = Pin> + Send>>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, _req: Request) -> Self::Future { + self.calls.lock().unwrap().push(self.name); + Box::pin(async { Ok(Response::new(tonic::body::Body::empty())) }) + } + } + + let calls = Arc::new(Mutex::new(Vec::new())); + let service = |name| RouteRecorder { + name, + calls: calls.clone(), + }; + let mut router = GrpcRouter::new(service("openshell"), service("reflection")); + + for path in [ + "/openshell.v1.OpenShell/Health", + "/grpc.reflection.v1.ServerReflection/ServerReflectionInfo", + ] { + router + .call( + Request::builder() + .uri(path) + .body(Empty::::new()) + .unwrap(), + ) + .await + .unwrap(); + } + + assert_eq!(*calls.lock().unwrap(), vec!["openshell", "reflection"]); + } + + #[tokio::test] + async fn reflection_protocol_serves_complete_public_descriptors_and_recovers_from_errors() { + use crate::auth::authenticator::test_support::MockAuthenticator; + use prost_reflect::{DescriptorPool, Value}; + use tonic_reflection::pb::v1::{ + ServerReflectionRequest, server_reflection_client::ServerReflectionClient, + server_reflection_request::MessageRequest, server_reflection_response::MessageResponse, + }; + + let reflection = + crate::reflection::build_gateway_reflection_service(&Config::new(None)).unwrap(); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let unrouted = tower::service_fn(|_request: Request| async { + Ok::<_, Infallible>(tonic::Status::unimplemented("test fallback").into_http()) + }); + let rejecting_oidc = Arc::new(MockAuthenticator::returning(Err( + tonic::Status::unauthenticated("OIDC credentials required"), + ))); + let grpc = AuthGrpcRouter::with_peer_identity( + GrpcRouter::new(unrouted, reflection), + Some(AuthenticatorChain::new(vec![rejecting_oidc])), + None, + None, + true, + false, + ); + let service = MultiplexedService::new(grpc, unrouted); + let server = tokio::spawn(async move { + loop { + let (stream, _) = listener.accept().await.unwrap(); + let service = service.clone(); + tokio::spawn(async move { + Builder::new(TokioExecutor::new()) + .serve_connection(TokioIo::new(stream), service) + .await + .unwrap(); + }); + } + }); + + let channel = tonic::transport::Channel::from_shared(format!("http://{addr}")) + .unwrap() + .connect() + .await + .unwrap(); + let mut client = ServerReflectionClient::new(channel); + let list_request = ServerReflectionRequest { + host: String::new(), + message_request: Some(MessageRequest::ListServices(String::new())), + }; + let descriptor_request = ServerReflectionRequest { + host: String::new(), + message_request: Some(MessageRequest::FileContainingSymbol( + "openshell.v1.OpenShell".to_string(), + )), + }; + let missing_request = ServerReflectionRequest { + host: String::new(), + message_request: Some(MessageRequest::FileContainingSymbol( + "openshell.v1.DoesNotExist".to_string(), + )), + }; + let list_after_error_request = ServerReflectionRequest { + host: String::new(), + message_request: Some(MessageRequest::ListServices(String::new())), + }; + let mut responses = client + .server_reflection_info(tokio_stream::iter([ + list_request, + descriptor_request, + missing_request, + list_after_error_request, + ])) + .await + .unwrap() + .into_inner(); + let response = responses.message().await.unwrap().unwrap(); + let Some(MessageResponse::ListServicesResponse(response)) = response.message_response + else { + panic!("expected a reflection list-services response"); + }; + let mut names: Vec<_> = response + .service + .into_iter() + .map(|service| service.name) + .collect(); + names.sort(); + + assert_eq!(names, vec!["openshell.v1.OpenShell"]); + + let descriptor_response = responses.message().await.unwrap().unwrap(); + let Some(MessageResponse::FileDescriptorResponse(response)) = + descriptor_response.message_response + else { + panic!("expected a reflection file-descriptor response"); + }; + let mut reflected_pool = DescriptorPool::new(); + for descriptor in &response.file_descriptor_proto { + reflected_pool + .decode_file_descriptor_proto(descriptor.as_slice()) + .unwrap(); + } + assert!( + reflected_pool + .get_service_by_name("openshell.v1.OpenShell") + .is_some(), + "a fresh descriptor pool must resolve the advertised service" + ); + assert!( + reflected_pool + .get_message_by_name("openshell.v1.HealthRequest") + .is_some(), + "the response must include imported public message descriptors" + ); + + let source_pool = DescriptorPool::decode(openshell_core::FILE_DESCRIPTOR_SET).unwrap(); + let source_auth = source_pool + .get_extension_by_name("openshell.options.v1.authorization") + .unwrap(); + let reflected_auth = reflected_pool + .get_extension_by_name("openshell.options.v1.authorization") + .unwrap(); + let source_health = source_pool + .get_service_by_name("openshell.v1.OpenShell") + .unwrap() + .methods() + .find(|method| method.name() == "Health") + .unwrap(); + let reflected_health = reflected_pool + .get_service_by_name("openshell.v1.OpenShell") + .unwrap() + .methods() + .find(|method| method.name() == "Health") + .unwrap(); + let source_options = source_health.options(); + let reflected_options = reflected_health.options(); + let Value::Message(source_auth_value) = &*source_options.get_extension(&source_auth) else { + panic!("source authorization option must be a message"); + }; + let Value::Message(reflected_auth_value) = + &*reflected_options.get_extension(&reflected_auth) + else { + panic!("reflected authorization option must be a message"); + }; + assert_eq!( + source_auth_value.encode_to_vec(), + reflected_auth_value.encode_to_vec() + ); + + let source_secret = source_pool + .get_extension_by_name("openshell.options.v1.secret") + .unwrap(); + let reflected_secret = reflected_pool + .get_extension_by_name("openshell.options.v1.secret") + .unwrap(); + let source_field = source_pool + .get_message_by_name("openshell.v1.TcpForwardInit") + .unwrap() + .get_field_by_name("authorization_token") + .unwrap(); + let reflected_field = reflected_pool + .get_message_by_name("openshell.v1.TcpForwardInit") + .unwrap() + .get_field_by_name("authorization_token") + .unwrap(); + assert_eq!( + source_field.options().get_extension(&source_secret), + reflected_field.options().get_extension(&reflected_secret), + "secret-field annotations must retain their original value" + ); + + let error_response = responses.message().await.unwrap().unwrap(); + let Some(MessageResponse::ErrorResponse(error)) = error_response.message_response else { + panic!("expected an in-band reflection error response"); + }; + assert_eq!(error.error_code, tonic::Code::NotFound as i32); + + let response_after_error = responses.message().await.unwrap().unwrap(); + assert!(matches!( + response_after_error.message_response, + Some(MessageResponse::ListServicesResponse(_)) + )); + server.abort(); + } + + #[tokio::test] + async fn reflection_rate_limit_charges_each_query_on_one_stream() { + use tonic::Code; + use tonic_reflection::pb::v1::{ + ServerReflectionRequest, server_reflection_client::ServerReflectionClient, + server_reflection_request::MessageRequest, + }; + + let config = Config::new(None).with_grpc_rate_limit(Some(1), Some(60)); + let reflection = crate::reflection::build_gateway_reflection_service(&config).unwrap(); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let unrouted = tower::service_fn(|_request: Request| async { + Ok::<_, Infallible>(tonic::Status::unimplemented("test fallback").into_http()) + }); + let service = MultiplexedService::new(GrpcRouter::new(unrouted, reflection), unrouted); + let server = tokio::spawn(async move { + loop { + let (stream, _) = listener.accept().await.unwrap(); + let service = service.clone(); + tokio::spawn(async move { + Builder::new(TokioExecutor::new()) + .serve_connection(TokioIo::new(stream), service) + .await + .unwrap(); + }); + } + }); + + let channel = tonic::transport::Channel::from_shared(format!("http://{addr}")) + .unwrap() + .connect() + .await + .unwrap(); + let mut client = ServerReflectionClient::new(channel); + let query = ServerReflectionRequest { + host: String::new(), + message_request: Some(MessageRequest::ListServices(String::new())), + }; + let mut responses = client + .server_reflection_info(tokio_stream::iter([query.clone(), query])) + .await + .unwrap() + .into_inner(); + + assert!(responses.message().await.unwrap().is_some()); + let status = responses + .message() + .await + .expect_err("second query on the same stream must be rate limited"); + assert_eq!(status.code(), Code::ResourceExhausted); + server.abort(); + } + + #[test] + fn reflection_descriptor_excludes_internal_service_protos() { + let descriptors = crate::reflection::gateway_reflection_descriptors().unwrap(); + let names: std::collections::BTreeSet<_> = descriptors + .iter() + .map(prost_reflect::FileDescriptor::name) + .collect(); + + assert!(names.contains("openshell.proto")); + assert!(names.contains("sandbox.proto")); + assert!(!names.contains("compute_driver.proto")); + assert!(!names.contains("credential_driver.proto")); + assert!(!names.contains("gateway_interceptor.proto")); + assert!(!names.contains("supervisor_middleware.proto")); + } + #[test] fn normalize_ws_tunnel() { assert_eq!(normalize_http_path("/_ws_tunnel"), "/_ws_tunnel"); @@ -2739,6 +3093,31 @@ mod tests { assert_eq!(grpc_status(&res).as_deref(), Some("16")); } + #[tokio::test] + async fn reflection_bypasses_oidc_and_mtls_user_authentication() { + let oidc = Arc::new(MockAuthenticator::returning(Err( + tonic::Status::unauthenticated("OIDC credentials required"), + ))); + let chain = AuthenticatorChain::new(vec![oidc]); + let (recorder, seen) = PrincipalRecorder::new(); + let mut router = + AuthGrpcRouter::with_peer_identity(recorder, Some(chain), None, None, true, false); + + let res = router + .call(empty_request( + "/grpc.reflection.v1.ServerReflection/ServerReflectionInfo", + )) + .await + .unwrap(); + + assert_eq!(res.status(), 200); + assert_eq!(grpc_status(&res), None); + assert!( + seen.lock().unwrap().is_none(), + "reflection must not receive an authenticated user principal" + ); + } + #[tokio::test] async fn unauthenticated_dev_user_fills_missing_principal_when_enabled() { let mock = Arc::new(MockAuthenticator::returning(Ok(None))); diff --git a/crates/openshell-server/src/reflection.rs b/crates/openshell-server/src/reflection.rs new file mode 100644 index 0000000000..2e22bb4a5b --- /dev/null +++ b/crates/openshell-server/src/reflection.rs @@ -0,0 +1,369 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Public gateway gRPC reflection service. + +use std::collections::{BTreeSet, HashMap}; +use std::sync::Arc; + +use prost_reflect::{DescriptorError, DescriptorPool, FileDescriptor}; +use prost_types::{DescriptorProto, EnumDescriptorProto, FileDescriptorProto}; +use tokio::sync::mpsc; +use tokio_stream::{Stream, StreamExt}; +use tonic::{Request, Response, Status, Streaming}; +use tonic_reflection::pb::v1::server_reflection_request::MessageRequest; +use tonic_reflection::pb::v1::server_reflection_response::MessageResponse; +use tonic_reflection::pb::v1::server_reflection_server::{ + ServerReflection, ServerReflectionServer, +}; +use tonic_reflection::pb::v1::{ + ErrorResponse, ExtensionNumberResponse, FileDescriptorResponse, ListServiceResponse, + ServerReflectionRequest, ServerReflectionResponse, ServiceResponse, +}; + +use openshell_core::Config; + +use crate::multiplex::GrpcRateLimiter; + +const REFLECTED_PROTO_ROOTS: &[&str] = &["openshell.proto"]; +const ADVERTISED_SERVICES: &[&str] = &["openshell.v1.OpenShell"]; + +pub type GatewayReflectionServer = ServerReflectionServer; + +/// Decode and filter the compiled descriptors to the public gateway schema. +pub fn gateway_reflection_descriptors() -> Result, DescriptorError> { + let descriptor_pool = DescriptorPool::decode(openshell_core::FILE_DESCRIPTOR_SET)?; + let mut included: BTreeSet = REFLECTED_PROTO_ROOTS + .iter() + .map(|name| (*name).to_string()) + .collect(); + + loop { + let before = included.len(); + for file in descriptor_pool.files() { + if included.contains(file.name()) { + included.extend( + file.dependencies() + .map(|dependency| dependency.name().to_string()), + ); + } + } + if included.len() == before { + break; + } + } + + Ok(descriptor_pool + .files() + .filter(|file| included.contains(file.name())) + .collect()) +} + +/// Build the immutable reflection index once at gateway service startup. +pub fn build_gateway_reflection_service( + config: &Config, +) -> Result { + let mut descriptors = gateway_reflection_descriptors()?; + descriptors + .extend(DescriptorPool::decode(tonic_reflection::pb::v1::FILE_DESCRIPTOR_SET)?.files()); + + let state = ReflectionState::new(descriptors); + Ok(ServerReflectionServer::new(GatewayReflectionService { + state: Arc::new(state), + limiter: GrpcRateLimiter::from_config(config), + })) +} + +#[derive(Debug)] +struct ReflectionState { + files: HashMap>, + symbols: HashMap>, +} + +#[derive(Debug)] +struct ReflectionFile { + descriptor: FileDescriptorProto, + encoded: Vec, +} + +impl ReflectionState { + fn new(descriptors: Vec) -> Self { + let mut state = Self { + files: HashMap::new(), + symbols: HashMap::new(), + }; + for descriptor in descriptors { + let name = descriptor.name().to_string(); + let descriptor = Arc::new(ReflectionFile { + descriptor: descriptor.file_descriptor_proto().clone(), + encoded: descriptor.encode_to_vec(), + }); + state.process_file(descriptor.clone()); + state.files.insert(name, descriptor); + } + state + } + + fn process_file(&mut self, file: Arc) { + let prefix = file + .descriptor + .package + .as_deref() + .unwrap_or_default() + .to_string(); + for message in &file.descriptor.message_type { + self.process_message(file.clone(), &prefix, message); + } + for enumeration in &file.descriptor.enum_type { + self.process_enum(file.clone(), &prefix, enumeration); + } + for service in &file.descriptor.service { + let Some(name) = service.name.as_deref() else { + continue; + }; + let service_name = qualified_name(&prefix, name); + self.symbols.insert(service_name.clone(), file.clone()); + for method in &service.method { + if let Some(name) = method.name.as_deref() { + self.symbols + .insert(qualified_name(&service_name, name), file.clone()); + } + } + } + } + + fn process_message( + &mut self, + file: Arc, + prefix: &str, + message: &DescriptorProto, + ) { + let Some(name) = message.name.as_deref() else { + return; + }; + let message_name = qualified_name(prefix, name); + self.symbols.insert(message_name.clone(), file.clone()); + for nested in &message.nested_type { + self.process_message(file.clone(), &message_name, nested); + } + for enumeration in &message.enum_type { + self.process_enum(file.clone(), &message_name, enumeration); + } + for field in &message.field { + if let Some(name) = field.name.as_deref() { + self.symbols + .insert(qualified_name(&message_name, name), file.clone()); + } + } + for oneof in &message.oneof_decl { + if let Some(name) = oneof.name.as_deref() { + self.symbols + .insert(qualified_name(&message_name, name), file.clone()); + } + } + } + + fn process_enum( + &mut self, + file: Arc, + prefix: &str, + enumeration: &EnumDescriptorProto, + ) { + let Some(name) = enumeration.name.as_deref() else { + return; + }; + let enum_name = qualified_name(prefix, name); + self.symbols.insert(enum_name.clone(), file.clone()); + for value in &enumeration.value { + if let Some(name) = value.name.as_deref() { + self.symbols + .insert(qualified_name(&enum_name, name), file.clone()); + } + } + } + + fn file_by_name(&self, name: &str) -> Result, Status> { + self.files + .get(name) + .cloned() + .ok_or_else(|| Status::not_found(format!("file '{name}' not found"))) + } + + fn file_by_symbol(&self, symbol: &str) -> Result, Status> { + self.symbols + .get(symbol) + .cloned() + .ok_or_else(|| Status::not_found(format!("symbol '{symbol}' not found"))) + } + + fn file_with_dependencies( + &self, + file: Arc, + sent: &mut BTreeSet, + ) -> Result>, Status> { + let mut descriptors = Vec::new(); + self.collect_file_with_dependencies(file, sent, &mut descriptors)?; + Ok(descriptors) + } + + fn collect_file_with_dependencies( + &self, + file: Arc, + sent: &mut BTreeSet, + descriptors: &mut Vec>, + ) -> Result<(), Status> { + let name = file + .descriptor + .name + .as_deref() + .ok_or_else(|| Status::internal("reflection descriptor is missing its filename"))?; + if !sent.insert(name.to_string()) { + return Ok(()); + } + + for dependency in &file.descriptor.dependency { + let dependency = self.files.get(dependency).cloned().ok_or_else(|| { + Status::internal(format!( + "reflection descriptor '{name}' has unavailable dependency '{dependency}'" + )) + })?; + self.collect_file_with_dependencies(dependency, sent, descriptors)?; + } + descriptors.push(file.encoded.clone()); + Ok(()) + } +} + +fn qualified_name(prefix: &str, name: &str) -> String { + if prefix.is_empty() { + name.to_string() + } else { + format!("{prefix}.{name}") + } +} + +/// Reflection implementation with a quota charged for every stream message. +#[derive(Clone, Debug)] +pub struct GatewayReflectionService { + state: Arc, + limiter: Option, +} + +#[tonic::async_trait] +impl ServerReflection for GatewayReflectionService { + type ServerReflectionInfoStream = ReflectionResponseStream; + + async fn server_reflection_info( + &self, + request: Request>, + ) -> Result, Status> { + let mut requests = request.into_inner(); + let (responses_tx, responses_rx) = mpsc::channel(1); + let state = self.state.clone(); + let limiter = self.limiter.clone(); + + tokio::spawn(async move { + let mut sent_descriptors = BTreeSet::new(); + while let Some(request) = requests.next().await { + let Ok(request) = request else { + return; + }; + if limiter.as_ref().is_some_and(|limiter| !limiter.allow()) { + let _ = responses_tx + .send(Err(Status::resource_exhausted( + "gRPC reflection query rate limit exceeded", + ))) + .await; + return; + } + + let response = match request.message_request.as_ref() { + Some(MessageRequest::FileByFilename(name)) => state + .file_by_name(name) + .and_then(|descriptor| { + state.file_with_dependencies(descriptor, &mut sent_descriptors) + }) + .map(|descriptors| { + MessageResponse::FileDescriptorResponse(FileDescriptorResponse { + file_descriptor_proto: descriptors, + }) + }), + Some(MessageRequest::FileContainingSymbol(symbol)) => state + .file_by_symbol(symbol) + .and_then(|descriptor| { + state.file_with_dependencies(descriptor, &mut sent_descriptors) + }) + .map(|descriptors| { + MessageResponse::FileDescriptorResponse(FileDescriptorResponse { + file_descriptor_proto: descriptors, + }) + }), + Some(MessageRequest::FileContainingExtension(_)) => { + Err(Status::not_found("extensions are not supported")) + } + Some(MessageRequest::AllExtensionNumbersOfType(_)) => { + Ok(MessageResponse::AllExtensionNumbersResponse( + ExtensionNumberResponse::default(), + )) + } + Some(MessageRequest::ListServices(_)) => { + Ok(MessageResponse::ListServicesResponse(ListServiceResponse { + service: ADVERTISED_SERVICES + .iter() + .map(|name| ServiceResponse { + name: (*name).to_string(), + }) + .collect(), + })) + } + None => Err(Status::invalid_argument("invalid MessageRequest")), + }; + + match response { + Ok(message_response) => { + let response = ServerReflectionResponse { + valid_host: request.host.clone(), + original_request: Some(request), + message_response: Some(message_response), + }; + if responses_tx.send(Ok(response)).await.is_err() { + return; + } + } + Err(status) => { + let response = ServerReflectionResponse { + valid_host: request.host.clone(), + original_request: Some(request), + message_response: Some(MessageResponse::ErrorResponse(ErrorResponse { + error_code: status.code() as i32, + error_message: status.message().to_string(), + })), + }; + if responses_tx.send(Ok(response)).await.is_err() { + return; + } + } + } + } + }); + + Ok(Response::new(ReflectionResponseStream { + inner: tokio_stream::wrappers::ReceiverStream::new(responses_rx), + })) + } +} + +pub struct ReflectionResponseStream { + inner: tokio_stream::wrappers::ReceiverStream>, +} + +impl Stream for ReflectionResponseStream { + type Item = Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::pin::Pin::new(&mut self.inner).poll_next(cx) + } +} diff --git a/docs/how-it-works/gateways/authentication.mdx b/docs/how-it-works/gateways/authentication.mdx index 3dbf6e524f..162a2caedb 100644 --- a/docs/how-it-works/gateways/authentication.mdx +++ b/docs/how-it-works/gateways/authentication.mdx @@ -69,6 +69,39 @@ The connection flow: 5. When mTLS user authentication is enabled, the gateway maps the verified certificate subject to a user principal. 6. The gateway authorizes the gRPC method. +### Inspect the API with grpcurl + +The primary gateway listener serves the gRPC reflection v1 protocol. Reflection +does not require application authentication, but the listener's TLS and client +certificate requirements still apply. + +For a local plaintext development gateway: + +```shell +grpcurl -plaintext localhost:18080 list +grpcurl -plaintext localhost:18080 describe openshell.v1.OpenShell +grpcurl -plaintext -d '{}' localhost:18080 openshell.v1.OpenShell/Health +``` + +For an mTLS gateway, use the bundle associated with the gateway: + +```shell +grpcurl \ + -cacert ~/.config/openshell/gateways//mtls/ca.crt \ + -cert ~/.config/openshell/gateways//mtls/tls.crt \ + -key ~/.config/openshell/gateways//mtls/tls.key \ + : list +``` + +Reflection advertises only `openshell.v1.OpenShell`. It does not advertise +internal driver, interceptor, or middleware services. + +When the gateway gRPC rate limit is enabled, reflection uses a separate counter +with the same request count and window. Opening the reflection RPC consumes one +gateway-wide request, and every query sent through that stream consumes one +reflection counter unit. Exceeding the reflection quota closes the stream with +`RESOURCE_EXHAUSTED`. + ### OIDC Gateways can validate OpenID Connect access tokens on gRPC requests. Configure OIDC when you want users, operators, or automation to authenticate with an identity provider such as Keycloak, Entra ID, or Okta. diff --git a/docs/how-it-works/gateways/configuration.mdx b/docs/how-it-works/gateways/configuration.mdx index 1308871cea..b1ce4c535f 100644 --- a/docs/how-it-works/gateways/configuration.mdx +++ b/docs/how-it-works/gateways/configuration.mdx @@ -165,6 +165,7 @@ guest_tls_cert = "/etc/openshell/certs/client.pem" guest_tls_key = "/etc/openshell/certs/client-key.pem" # Optional gRPC rate limit. Both values must be positive to enable the limit. +# Reflection queries use a separate counter with the same count and window. # Set either value to 0, or omit both, to disable rate limiting. grpc_rate_limit_requests = 120 grpc_rate_limit_window_seconds = 60