From dd83be542323fe0e7bcd90e52f383ba868c9b863 Mon Sep 17 00:00:00 2001 From: Krzysztof Malczuk Date: Thu, 3 Sep 2026 14:38:09 +0100 Subject: [PATCH 1/2] feat(gateway): add gRPC server reflection Signed-off-by: Krzysztof Malczuk --- Cargo.lock | 15 + Cargo.toml | 1 + README.md | 17 ++ crates/openshell-server/Cargo.toml | 1 + crates/openshell-server/src/auth/oidc.rs | 10 +- crates/openshell-server/src/multiplex.rs | 260 +++++++++++++++++- docs/how-it-works/gateways/authentication.mdx | 28 ++ 7 files changed, 328 insertions(+), 4 deletions(-) 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/multiplex.rs b/crates/openshell-server/src/multiplex.rs index 8988542b1a..2b7323925c 100644 --- a/crates/openshell-server/src/multiplex.rs +++ b/crates/openshell-server/src/multiplex.rs @@ -28,6 +28,7 @@ use opentelemetry::propagation::TextMapPropagator; use opentelemetry::trace::TraceContextExt as _; use opentelemetry_sdk::propagation::TraceContextPropagator; use prost::Message; +use prost_types::FileDescriptorSet; use std::collections::BTreeMap; use std::convert::Infallible; use std::future::Future; @@ -199,6 +200,39 @@ 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; +const REFLECTED_PROTO_ROOTS: &[&str] = &["openshell.proto"]; + +/// Restrict reflection to the public gateway APIs and their imported types. +fn gateway_reflection_descriptor_set() -> Result { + let mut descriptor_set = FileDescriptorSet::decode(openshell_core::FILE_DESCRIPTOR_SET)?; + let mut included: std::collections::BTreeSet = REFLECTED_PROTO_ROOTS + .iter() + .map(|name| (*name).to_string()) + .collect(); + + loop { + let before = included.len(); + for file in &descriptor_set.file { + if file + .name + .as_ref() + .is_some_and(|name| included.contains(name)) + { + included.extend(file.dependency.iter().cloned()); + } + } + if included.len() == before { + break; + } + } + + descriptor_set.file.retain(|file| { + file.name + .as_ref() + .is_some_and(|name| included.contains(name)) + }); + Ok(descriptor_set) +} /// Concurrent HTTP/2 streams allowed per connection. Sits above the /// per-replica pending relay budget so pooled peer connections are bounded by @@ -243,6 +277,10 @@ impl MultiplexService { self.state.gateway_interceptors.clone(), Some(self.state.clone()), ); + let reflection = tonic_reflection::server::Builder::configure() + .register_file_descriptor_set(gateway_reflection_descriptor_set()?) + .with_service_name("openshell.v1.OpenShell") + .build_v1()?; 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 +288,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 @@ -857,6 +895,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 +2581,151 @@ 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 running_primary_gateway_reflection_advertises_only_public_services() { + use crate::auth::authenticator::test_support::MockAuthenticator; + use tonic_reflection::pb::v1::{ + ServerReflectionRequest, server_reflection_client::ServerReflectionClient, + server_reflection_request::MessageRequest, server_reflection_response::MessageResponse, + }; + + let reflection = tonic_reflection::server::Builder::configure() + .register_file_descriptor_set(gateway_reflection_descriptor_set().unwrap()) + .with_service_name("openshell.v1.OpenShell") + .build_v1() + .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 = GatewayListenerContextService::new( + MultiplexedService::new(grpc, unrouted), + GatewayListenerScope::Primary, + ); + 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 request = ServerReflectionRequest { + host: String::new(), + message_request: Some(MessageRequest::ListServices(String::new())), + }; + let response = client + .server_reflection_info(tokio_stream::iter([request])) + .await + .unwrap() + .into_inner() + .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"]); + server.abort(); + } + + #[test] + fn reflection_descriptor_excludes_internal_service_protos() { + let descriptors = gateway_reflection_descriptor_set().unwrap(); + let names: std::collections::BTreeSet<_> = descriptors + .file + .iter() + .filter_map(|file| file.name.as_deref()) + .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 +2972,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/docs/how-it-works/gateways/authentication.mdx b/docs/how-it-works/gateways/authentication.mdx index 3dbf6e524f..f53970c7e3 100644 --- a/docs/how-it-works/gateways/authentication.mdx +++ b/docs/how-it-works/gateways/authentication.mdx @@ -69,6 +69,34 @@ 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. Callback-only +compute-driver listeners do not serve reflection. + ### 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. From 799c99b766a54a4434945663b9d3df78ec58b8d8 Mon Sep 17 00:00:00 2001 From: Krzysztof Malczuk Date: Tue, 22 Sep 2026 15:51:12 +0100 Subject: [PATCH 2/2] fix(gateway): rate limit gRPC reflection queries Signed-off-by: Krzysztof Malczuk --- crates/openshell-server/src/lib.rs | 7 + crates/openshell-server/src/multiplex.rs | 141 ++++---- crates/openshell-server/src/reflection.rs | 315 ++++++++++++++++++ docs/how-it-works/gateways/authentication.mdx | 9 +- docs/how-it-works/gateways/configuration.mdx | 1 + 5 files changed, 413 insertions(+), 60 deletions(-) create mode 100644 crates/openshell-server/src/reflection.rs diff --git a/crates/openshell-server/src/lib.rs b/crates/openshell-server/src/lib.rs index d6d35f6ba3..2d07bc73bb 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(), diff --git a/crates/openshell-server/src/multiplex.rs b/crates/openshell-server/src/multiplex.rs index 2b7323925c..7d2ba6569f 100644 --- a/crates/openshell-server/src/multiplex.rs +++ b/crates/openshell-server/src/multiplex.rs @@ -28,7 +28,6 @@ use opentelemetry::propagation::TextMapPropagator; use opentelemetry::trace::TraceContextExt as _; use opentelemetry_sdk::propagation::TraceContextPropagator; use prost::Message; -use prost_types::FileDescriptorSet; use std::collections::BTreeMap; use std::convert::Infallible; use std::future::Future; @@ -200,40 +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; -const REFLECTED_PROTO_ROOTS: &[&str] = &["openshell.proto"]; - -/// Restrict reflection to the public gateway APIs and their imported types. -fn gateway_reflection_descriptor_set() -> Result { - let mut descriptor_set = FileDescriptorSet::decode(openshell_core::FILE_DESCRIPTOR_SET)?; - let mut included: std::collections::BTreeSet = REFLECTED_PROTO_ROOTS - .iter() - .map(|name| (*name).to_string()) - .collect(); - - loop { - let before = included.len(); - for file in &descriptor_set.file { - if file - .name - .as_ref() - .is_some_and(|name| included.contains(name)) - { - included.extend(file.dependency.iter().cloned()); - } - } - if included.len() == before { - break; - } - } - - descriptor_set.file.retain(|file| { - file.name - .as_ref() - .is_some_and(|name| included.contains(name)) - }); - Ok(descriptor_set) -} - /// 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. @@ -277,10 +242,7 @@ impl MultiplexService { self.state.gateway_interceptors.clone(), Some(self.state.clone()), ); - let reflection = tonic_reflection::server::Builder::configure() - .register_file_descriptor_set(gateway_reflection_descriptor_set()?) - .with_service_name("openshell.v1.OpenShell") - .build_v1()?; + 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(), @@ -774,7 +736,7 @@ impl GrpcRateLimiter { }) } - fn allow(&self) -> bool { + pub(crate) fn allow(&self) -> bool { let now = Instant::now(); let mut state = self .state @@ -2637,11 +2599,8 @@ mod tests { server_reflection_request::MessageRequest, server_reflection_response::MessageResponse, }; - let reflection = tonic_reflection::server::Builder::configure() - .register_file_descriptor_set(gateway_reflection_descriptor_set().unwrap()) - .with_service_name("openshell.v1.OpenShell") - .build_v1() - .unwrap(); + 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 { @@ -2658,10 +2617,7 @@ mod tests { true, false, ); - let service = GatewayListenerContextService::new( - MultiplexedService::new(grpc, unrouted), - GatewayListenerScope::Primary, - ); + let service = MultiplexedService::new(grpc, unrouted); let server = tokio::spawn(async move { loop { let (stream, _) = listener.accept().await.unwrap(); @@ -2681,19 +2637,22 @@ mod tests { .await .unwrap(); let mut client = ServerReflectionClient::new(channel); - let request = ServerReflectionRequest { + let list_request = ServerReflectionRequest { host: String::new(), message_request: Some(MessageRequest::ListServices(String::new())), }; - let response = client - .server_reflection_info(tokio_stream::iter([request])) - .await - .unwrap() - .into_inner() - .message() + let descriptor_request = ServerReflectionRequest { + host: String::new(), + message_request: Some(MessageRequest::FileContainingSymbol( + "openshell.v1.OpenShell".to_string(), + )), + }; + let mut responses = client + .server_reflection_info(tokio_stream::iter([list_request, descriptor_request])) .await .unwrap() - .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"); @@ -2706,12 +2665,78 @@ mod tests { 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 descriptor = prost_types::FileDescriptorProto::decode( + response.file_descriptor_proto.first().unwrap().as_slice(), + ) + .unwrap(); + assert_eq!(descriptor.name.as_deref(), Some("openshell.proto")); + 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 = gateway_reflection_descriptor_set().unwrap(); + let descriptors = crate::reflection::gateway_reflection_descriptor_set().unwrap(); let names: std::collections::BTreeSet<_> = descriptors .file .iter() diff --git a/crates/openshell-server/src/reflection.rs b/crates/openshell-server/src/reflection.rs new file mode 100644 index 0000000000..a8eb4f38cb --- /dev/null +++ b/crates/openshell-server/src/reflection.rs @@ -0,0 +1,315 @@ +// 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::Message; +use prost_types::{DescriptorProto, EnumDescriptorProto, FileDescriptorProto, FileDescriptorSet}; +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::{ + 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_descriptor_set() -> Result { + let mut descriptor_set = FileDescriptorSet::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_set.file { + if file + .name + .as_ref() + .is_some_and(|name| included.contains(name)) + { + included.extend(file.dependency.iter().cloned()); + } + } + if included.len() == before { + break; + } + } + + descriptor_set.file.retain(|file| { + file.name + .as_ref() + .is_some_and(|name| included.contains(name)) + }); + Ok(descriptor_set) +} + +/// Build the immutable reflection index once at gateway service startup. +pub fn build_gateway_reflection_service( + config: &Config, +) -> Result { + let mut descriptors = gateway_reflection_descriptor_set()?; + descriptors + .file + .extend(FileDescriptorSet::decode(tonic_reflection::pb::v1::FILE_DESCRIPTOR_SET)?.file); + + 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>, +} + +impl ReflectionState { + fn new(descriptors: FileDescriptorSet) -> Self { + let mut state = Self { + files: HashMap::new(), + symbols: HashMap::new(), + }; + for descriptor in descriptors.file { + let Some(name) = descriptor.name.clone() else { + continue; + }; + let descriptor = Arc::new(descriptor); + state.process_file(descriptor.clone()); + state.files.insert(name, descriptor); + } + state + } + + fn process_file(&mut self, file: Arc) { + let prefix = file.package.as_deref().unwrap_or_default().to_string(); + for message in &file.message_type { + self.process_message(file.clone(), &prefix, message); + } + for enumeration in &file.enum_type { + self.process_enum(file.clone(), &prefix, enumeration); + } + for service in &file.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 encode_file(file: &FileDescriptorProto) -> Result, Status> { + let mut encoded = Vec::new(); + file.encode(&mut encoded) + .map_err(|_| Status::internal("failed to encode reflection descriptor"))?; + Ok(encoded) + } + + fn file_by_name(&self, name: &str) -> Result, Status> { + self.files.get(name).map_or_else( + || Err(Status::not_found(format!("file '{name}' not found"))), + |file| Self::encode_file(file), + ) + } + + fn file_by_symbol(&self, symbol: &str) -> Result, Status> { + self.symbols.get(symbol).map_or_else( + || Err(Status::not_found(format!("symbol '{symbol}' not found"))), + |file| Self::encode_file(file), + ) + } +} + +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 { + 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).map(|descriptor| { + MessageResponse::FileDescriptorResponse(FileDescriptorResponse { + file_descriptor_proto: vec![descriptor], + }) + }) + } + Some(MessageRequest::FileContainingSymbol(symbol)) => { + state.file_by_symbol(symbol).map(|descriptor| { + MessageResponse::FileDescriptorResponse(FileDescriptorResponse { + file_descriptor_proto: vec![descriptor], + }) + }) + } + 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 _ = responses_tx.send(Err(status)).await; + 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 f53970c7e3..162a2caedb 100644 --- a/docs/how-it-works/gateways/authentication.mdx +++ b/docs/how-it-works/gateways/authentication.mdx @@ -94,8 +94,13 @@ grpcurl \ ``` Reflection advertises only `openshell.v1.OpenShell`. It does not advertise -internal driver, interceptor, or middleware services. Callback-only -compute-driver listeners do not serve reflection. +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 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