Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
81 changes: 61 additions & 20 deletions engine/packages/guard-core/src/proxy_service.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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<BodyIncoming>) -> Result<Response<ResponseBody>> {
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();

Expand All @@ -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)
Expand Down Expand Up @@ -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,
Expand All @@ -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(),
Expand Down Expand Up @@ -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 {
Expand All @@ -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();
Expand Down Expand Up @@ -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",
);
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
{
Expand Down Expand Up @@ -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;
}
Expand Down Expand Up @@ -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?;

Expand Down Expand Up @@ -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?;

Expand All @@ -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?;

Expand Down Expand Up @@ -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 {
Expand All @@ -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(),
Expand Down
17 changes: 17 additions & 0 deletions engine/packages/guard-core/src/request_context.rs
Original file line number Diff line number Diff line change
@@ -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},
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -60,6 +64,7 @@ impl RequestContext {
RequestContext {
remote_addr,
ray_id,
external_ray_id,
req_id,
host,
hostname,
Expand Down Expand Up @@ -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<String, String>) {
headers.insert(
X_RIVET_RAY_ID.as_str().to_owned(),
self.external_ray_id.clone(),
);
}

pub fn req_id(&self) -> Id {
self.req_id
}
Expand Down
3 changes: 1 addition & 2 deletions engine/packages/guard-core/src/utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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();
Expand Down
63 changes: 62 additions & 1 deletion engine/packages/guard-core/tests/proxy.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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::<bytes::Bytes>::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();
Expand All @@ -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::<bytes::Bytes>::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::<bytes::Bytes>::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]
Expand Down
Loading
Loading