diff --git a/engine/packages/guard-core/src/proxy_service.rs b/engine/packages/guard-core/src/proxy_service.rs index 4aba08aaf0..de4f82b1e4 100644 --- a/engine/packages/guard-core/src/proxy_service.rs +++ b/engine/packages/guard-core/src/proxy_service.rs @@ -16,7 +16,6 @@ use rand::seq::SliceRandom; use rivet_api_builder::{RequestIds, X_RIVET_RAY_ID}; use rivet_error::RivetError; use rivet_metrics::GaugeGuardExt; -use rivet_util::Id; use tracing_opentelemetry::OpenTelemetrySpanExt; use rivet_runner_protocol as protocol; @@ -44,6 +43,23 @@ pub const X_FORWARDED_FOR: HeaderName = HeaderName::from_static("x-forwarded-for pub const X_RIVET_ERROR: HeaderName = HeaderName::from_static("x-rivet-error"); const WEBSOCKET_CLOSE_LINGER: Duration = Duration::from_millis(5); // Keep TCP connection open briefly after WebSocket close +const MAX_EXTERNAL_RAY_ID_LEN: usize = 30; + +/// Returns `value` when it is a ray ID actors accept: 1 to 30 characters of +/// `[A-Za-z0-9_-]`. RivetKit enforces the same bound on its side. The engine +/// cannot depend on RivetKit crates, so the rule is repeated here. +fn bounded_external_ray_id(value: &str) -> Option<&str> { + if !value.is_empty() + && value.len() <= MAX_EXTERNAL_RAY_ID_LEN + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-')) + { + Some(value) + } else { + None + } +} fn websocket_config(guard_config: &rivet_config::config::guard::Guard) -> WebSocketConfig { WebSocketConfig::default() @@ -356,12 +372,22 @@ impl ProxyService { } /// Process an individual request. - #[tracing::instrument(name = "guard_request", skip_all, fields(ray_id, req_id, uri=%utils::redact_uri_for_logs(req.uri())))] + #[tracing::instrument(name = "guard_request", skip_all, fields(ray_id, external_ray_id, req_id, uri=%utils::redact_uri_for_logs(req.uri())))] pub async fn process(&self, mut req: Request) -> Result> { let start_time = Instant::now(); let request_ids = RequestIds::new(self.state.config.dc_label()); req.extensions_mut().insert(request_ids); + let external_ray_id = req + .headers() + .get(X_RIVET_RAY_ID) + .and_then(|value| value.to_str().ok()) + .and_then(bounded_external_ray_id) + .map(str::to_owned) + .unwrap_or_else(|| request_ids.ray_id.to_string()); + if let Ok(value) = HeaderValue::from_str(&external_ray_id) { + req.headers_mut().insert(X_RIVET_RAY_ID, value); + } let current_span = tracing::Span::current(); @@ -384,13 +410,9 @@ impl ProxyService { current_span.record("req_id", request_ids.req_id.to_string()); current_span.record("ray_id", request_ids.ray_id.to_string()); + current_span.record("external_ray_id", external_ray_id.as_str()); // Extract request information for logging and analytics before consuming the request - let incoming_ray_id = req - .headers() - .get(X_RIVET_RAY_ID) - .and_then(|h| h.to_str().ok()) - .and_then(|id| Id::parse(id).ok()); let host = req .headers() .get(hyper::header::HOST) @@ -432,6 +454,7 @@ impl ProxyService { let mut req_ctx = RequestContext::new( self.remote_addr, request_ids.ray_id, + external_ray_id, request_ids.req_id, host, path, @@ -450,8 +473,8 @@ impl ProxyService { // Debug log request information with structured fields (Apache-like access log) tracing::debug!( - ?incoming_ray_id, - ray_id=?req_ctx.ray_id, + ray_id=%req_ctx.ray_id, + external_ray_id=%req_ctx.external_ray_id, req_id=?req_ctx.req_id, method=%req_ctx.method, path=%req_ctx.path_for_logs(), @@ -507,6 +530,7 @@ impl ProxyService { tracing::debug!("Client WebSocket upgrade for error proxy successful"); let active_guard = metrics::WEBSOCKET_ACTIVE.inc_guard(); + let external_ray_id = req_ctx.external_ray_id.clone(); self.state.tasks.spawn( async move { @@ -522,7 +546,7 @@ impl ProxyService { return; } }; - let frame = utils::err_to_close_frame(err, request_ids.ray_id); + let frame = utils::err_to_close_frame(err, &external_ray_id); // Manual conversion to handle different tungstenite versions let code_num: u16 = frame.code.into(); @@ -608,15 +632,18 @@ impl ProxyService { } // Add ray_id to response headers - if let Ok(ray_id_value) = request_ids.ray_id.to_string().parse() { + if let Ok(ray_id_value) = HeaderValue::from_str(req_ctx.external_ray_id()) { if let Some(existing_ray_id_value) = res .headers() .get(X_RIVET_RAY_ID) .and_then(|h| h.to_str().ok()) { - if ray_id_value != existing_ray_id_value { + // api-builder sets the guard's own id on its responses, which is expected. + if ray_id_value != existing_ray_id_value + && existing_ray_id_value != req_ctx.ray_id.to_string() + { tracing::warn!( - expected_ray_id=%request_ids.ray_id, + expected_ray_id=%req_ctx.external_ray_id, received_ray_id=%existing_ray_id_value, "downstream service set ray id header to a different value", ); @@ -687,8 +714,8 @@ impl ProxyService { // Log information about the completed request tracing::debug!( - ?incoming_ray_id, - ray_id=?req_ctx.ray_id, + ray_id=%req_ctx.ray_id, + external_ray_id=%req_ctx.external_ray_id, req_id=?req_ctx.req_id, method = %req_ctx.method, path = %req_ctx.path, @@ -1309,7 +1336,7 @@ impl ProxyService { match client_sink .send(utils::to_hyper_close(Some(utils::err_to_close_frame( err, - req_ctx.ray_id, + req_ctx.external_ray_id(), )))) .await { @@ -1373,7 +1400,10 @@ impl ProxyService { "websocket target changed to custom serve" ); let _ = client_ws - .close(Some(utils::err_to_close_frame(err, req_ctx.ray_id))) + .close(Some(utils::err_to_close_frame( + err, + req_ctx.external_ray_id(), + ))) .await; return; } @@ -1820,7 +1850,10 @@ impl ProxyService { // Close WebSocket with error ws_handle .send(utils::to_hyper_close(Some( - utils::err_to_close_frame(err, req_ctx.ray_id), + utils::err_to_close_frame( + err, + req_ctx.external_ray_id(), + ), ))) .await?; @@ -1865,7 +1898,10 @@ impl ProxyService { ); ws_handle .send(utils::to_hyper_close(Some( - utils::err_to_close_frame(err, req_ctx.ray_id), + utils::err_to_close_frame( + err, + req_ctx.external_ray_id(), + ), ))) .await?; @@ -1884,7 +1920,10 @@ impl ProxyService { ); ws_handle .send(utils::to_hyper_close(Some( - utils::err_to_close_frame(err, req_ctx.ray_id), + utils::err_to_close_frame( + err, + req_ctx.external_ray_id(), + ), ))) .await?; @@ -1973,6 +2012,7 @@ impl ProxyServiceFactory { #[cfg(test)] mod tests { use super::*; + use rivet_util::Id; fn test_state(guard: rivet_config::config::guard::Guard) -> ProxyState { let config = rivet_config::Config::from_root(rivet_config::config::Root { @@ -1990,6 +2030,7 @@ mod tests { RequestContext::new( "127.0.0.1:12345".parse().unwrap(), Id::v1(uuid::Uuid::nil(), 0), + "test-ray".to_owned(), Id::v1(uuid::Uuid::nil(), 1), "example.com".to_owned(), "/actors".to_owned(), diff --git a/engine/packages/guard-core/src/request_context.rs b/engine/packages/guard-core/src/request_context.rs index 76c01f66ce..18a3faf612 100644 --- a/engine/packages/guard-core/src/request_context.rs +++ b/engine/packages/guard-core/src/request_context.rs @@ -1,8 +1,10 @@ use anyhow::{Context, Result}; use hyper::{Method, header::HeaderMap}; +use rivet_api_builder::X_RIVET_RAY_ID; use rivet_runner_protocol as protocol; use rivet_util::Id; use std::{ + collections::HashMap, net::{IpAddr, SocketAddr}, sync::Arc, time::{Duration, Instant}, @@ -15,6 +17,7 @@ use crate::utils::InFlightPermit; pub struct RequestContext { pub(crate) remote_addr: SocketAddr, pub(crate) ray_id: Id, + pub(crate) external_ray_id: String, pub(crate) req_id: Id, /// Entire host including port (if present) pub(crate) host: String, @@ -44,6 +47,7 @@ impl RequestContext { pub(crate) fn new( remote_addr: SocketAddr, ray_id: Id, + external_ray_id: String, req_id: Id, host: String, path: String, @@ -60,6 +64,7 @@ impl RequestContext { RequestContext { remote_addr, ray_id, + external_ray_id, req_id, host, hostname, @@ -90,6 +95,18 @@ impl RequestContext { self.ray_id } + pub fn external_ray_id(&self) -> &str { + &self.external_ray_id + } + + /// Adds this request's validated external ray ID to actor headers. + pub fn forward_ray(&self, headers: &mut HashMap) { + headers.insert( + X_RIVET_RAY_ID.as_str().to_owned(), + self.external_ray_id.clone(), + ); + } + pub fn req_id(&self) -> Id { self.req_id } diff --git a/engine/packages/guard-core/src/utils.rs b/engine/packages/guard-core/src/utils.rs index 9ae667bc80..5ce81aca2b 100644 --- a/engine/packages/guard-core/src/utils.rs +++ b/engine/packages/guard-core/src/utils.rs @@ -9,7 +9,6 @@ use rivet_api_builder::{ErrorResponse, RawErrorResponse}; use rivet_error::{INTERNAL_ERROR, RivetError}; use rivet_metrics::{GaugeGuardExt, IntGaugeGuard}; use rivet_runner_protocol as protocol; -use rivet_util::Id; use rivet_util::throttle::{RateLimitMethod, RateLimiter}; use std::sync::Arc; use std::time::Duration; @@ -491,7 +490,7 @@ pub fn is_ws_hibernate(err: &anyhow::Error) -> bool { } } -pub(crate) fn err_to_close_frame(err: anyhow::Error, ray_id: Id) -> CloseFrame { +pub(crate) fn err_to_close_frame(err: anyhow::Error, ray_id: &str) -> CloseFrame { metrics::WEBSOCKET_CLOSE_ERROR_TOTAL .with_label_values(&[&error_metric_label(&err)]) .inc(); diff --git a/engine/packages/guard-core/tests/proxy.rs b/engine/packages/guard-core/tests/proxy.rs index c0ab59b107..768e3f0c3a 100644 --- a/engine/packages/guard-core/tests/proxy.rs +++ b/engine/packages/guard-core/tests/proxy.rs @@ -2,6 +2,7 @@ mod common; use bytes::Bytes; use hyper::{Method, StatusCode}; +use rivet_api_builder::X_RIVET_RAY_ID; use rivet_util::Id; use std::sync::Arc; use std::time::{Duration, Instant}; @@ -48,7 +49,7 @@ async fn test_basic_proxy_functionality() { } #[tokio::test] -async fn test_proxy_forwards_headers() { +async fn test_proxy_forwards_headers_and_uses_one_bounded_ray_id() { init_tracing(); // Set up a test server that echoes back headers @@ -90,11 +91,16 @@ async fn test_proxy_forwards_headers() { .header(hyper::header::HOST, "example.com") .header("X-Custom-Header", "test-value") .header("X-Another-Header", "another-value") + .header(X_RIVET_RAY_ID, "caller-ray_123") .body(http_body_util::Empty::::new()) .unwrap(); let response = client.request(request).await.unwrap(); assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response.headers().get(X_RIVET_RAY_ID).unwrap(), + "caller-ray_123" + ); // Check that our custom headers were forwarded let last_request = test_server.last_request().unwrap(); @@ -108,6 +114,61 @@ async fn test_proxy_forwards_headers() { last_request.headers.get("x-another-header").unwrap(), "another-value" ); + assert_eq!( + last_request.headers.get(X_RIVET_RAY_ID.as_str()).unwrap(), + "caller-ray_123" + ); + + let invalid_ray_id = "a".repeat(31); + let request = hyper::Request::builder() + .method(Method::GET) + .uri(format!("http://{}/echo", guard_addr)) + .header(hyper::header::HOST, "example.com") + .header(X_RIVET_RAY_ID, &invalid_ray_id) + .body(http_body_util::Empty::::new()) + .unwrap(); + let response = client.request(request).await.unwrap(); + let response_ray_id = response + .headers() + .get(X_RIVET_RAY_ID) + .unwrap() + .to_str() + .unwrap(); + assert_ne!(response_ray_id, invalid_ray_id); + assert!(Id::parse(response_ray_id).is_ok()); + assert_eq!( + test_server + .last_request() + .unwrap() + .headers + .get(X_RIVET_RAY_ID.as_str()) + .unwrap(), + response_ray_id + ); + + let request = hyper::Request::builder() + .method(Method::GET) + .uri(format!("http://{}/echo", guard_addr)) + .header(hyper::header::HOST, "example.com") + .body(http_body_util::Empty::::new()) + .unwrap(); + let response = client.request(request).await.unwrap(); + let response_ray_id = response + .headers() + .get(X_RIVET_RAY_ID) + .unwrap() + .to_str() + .unwrap(); + assert!(Id::parse(response_ray_id).is_ok()); + assert_eq!( + test_server + .last_request() + .unwrap() + .headers + .get(X_RIVET_RAY_ID.as_str()) + .unwrap(), + response_ray_id + ); } #[tokio::test] diff --git a/engine/packages/pegboard-gateway2/src/lib.rs b/engine/packages/pegboard-gateway2/src/lib.rs index 290ceec819..67ee28c7a9 100644 --- a/engine/packages/pegboard-gateway2/src/lib.rs +++ b/engine/packages/pegboard-gateway2/src/lib.rs @@ -127,7 +127,7 @@ impl PegboardGateway2 { let request_id = req_ctx.in_flight_request_id()?; // Extract request parts - let headers = req + let mut headers = req .headers() .iter() .filter_map(|(name, value)| { @@ -137,6 +137,7 @@ impl PegboardGateway2 { .map(|value_str| (name.to_string(), value_str.to_string())) }) .collect::>(); + req_ctx.forward_ray(&mut headers); // NOTE: Size constraints have already been applied by guard let body_bytes = req @@ -360,6 +361,7 @@ impl PegboardGateway2 { request_headers.insert(name.to_string(), value_str.to_string()); } } + req_ctx.forward_ray(&mut request_headers); let (mut stopped_sub, _) = tokio::try_join!( ctx.subscribe::(("actor_id", self.actor_id)), diff --git a/engine/packages/pegboard-gateway3/src/http_stream/handler.rs b/engine/packages/pegboard-gateway3/src/http_stream/handler.rs index 1f83b4e535..1489d54cfd 100644 --- a/engine/packages/pegboard-gateway3/src/http_stream/handler.rs +++ b/engine/packages/pegboard-gateway3/src/http_stream/handler.rs @@ -1,4 +1,5 @@ use std::{ + collections::HashMap, sync::{ Arc, atomic::{AtomicU64, Ordering}, @@ -71,7 +72,7 @@ impl PegboardGateway3 { req_ctx.request_body_is_end_stream(), ); let request_id = req_ctx.in_flight_request_id()?; - let headers = req_ctx + let mut headers: HashMap = req_ctx .headers() .iter() .filter_map(|(name, value)| { @@ -81,6 +82,7 @@ impl PegboardGateway3 { .map(|value| (name.to_string(), value.to_owned())) }) .collect(); + req_ctx.forward_ray(&mut headers); let (mut stopped_sub, _) = tokio::try_join!( ctx.subscribe::(("actor_id", self.actor_id)), pegboard::utils::ensure_ns_metrics_exporter_for_namespace(ctx, self.namespace_id), diff --git a/engine/packages/pegboard-gateway3/src/lib.rs b/engine/packages/pegboard-gateway3/src/lib.rs index 8fb7ada9da..b90daf101d 100644 --- a/engine/packages/pegboard-gateway3/src/lib.rs +++ b/engine/packages/pegboard-gateway3/src/lib.rs @@ -134,6 +134,7 @@ impl PegboardGateway3 { request_headers.insert(name.to_string(), value_str.to_string()); } } + req_ctx.forward_ray(&mut request_headers); let (mut stopped_sub, _) = tokio::try_join!( ctx.subscribe::(("actor_id", self.actor_id)),