mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 02:47:45 +08:00
refactor(workspace): enforce layered crate boundaries
This commit is contained in:
@@ -159,6 +159,15 @@ fn relay_header_timeout(meta: &protocol::RequestMeta) -> Duration {
|
||||
Duration::from_millis(resolve_tunnel_request_timeouts(meta).first_byte_ms)
|
||||
}
|
||||
|
||||
fn is_rollout_probe_request(headers: &HeaderMap, forwarded_by_gateway: bool) -> bool {
|
||||
forwarded_by_gateway
|
||||
&& headers
|
||||
.get(crate::tunnel::TUNNEL_RELAY_ROLLOUT_PROBE_HEADER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| value == crate::tunnel::TUNNEL_RELAY_ROLLOUT_PROBE_VALUE)
|
||||
}
|
||||
|
||||
pub async fn relay_request(
|
||||
Path(node_id): Path<String>,
|
||||
State(state): State<AppState>,
|
||||
@@ -171,6 +180,7 @@ pub async fn relay_request(
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty());
|
||||
let rollout_probe = is_rollout_probe_request(request.headers(), forwarded_by_gateway);
|
||||
if !addr.ip().is_loopback() && !forwarded_by_gateway {
|
||||
return tunnel_error_response(
|
||||
StatusCode::FORBIDDEN,
|
||||
@@ -348,12 +358,16 @@ pub async fn relay_request(
|
||||
);
|
||||
}
|
||||
};
|
||||
if let Err(error) = record_proxy_upgrade_traffic_success(state.data.as_ref(), &node_id).await {
|
||||
warn!(
|
||||
node_id = %node_id,
|
||||
error = %error,
|
||||
"failed to record proxy upgrade traffic confirmation"
|
||||
);
|
||||
if !rollout_probe {
|
||||
if let Err(error) =
|
||||
record_proxy_upgrade_traffic_success(state.data.as_ref(), &node_id).await
|
||||
{
|
||||
warn!(
|
||||
node_id = %node_id,
|
||||
error = %error,
|
||||
"failed to record proxy upgrade traffic confirmation"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let Some(mut body_rx) = stream.take_body_receiver() else {
|
||||
@@ -462,8 +476,8 @@ mod tests {
|
||||
use super::super::hub::ProxyConn;
|
||||
use super::super::{protocol, AppState, ConnConfig, ControlPlaneClient};
|
||||
use super::{
|
||||
relay_header_timeout, relay_request, Body, Request, SocketAddr, StatusCode,
|
||||
TUNNEL_ERROR_HEADER,
|
||||
is_rollout_probe_request, relay_header_timeout, relay_request, Body, HeaderMap, Request,
|
||||
SocketAddr, StatusCode, TUNNEL_ERROR_HEADER,
|
||||
};
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::maintenance::start_proxy_upgrade_rollout;
|
||||
@@ -482,6 +496,20 @@ mod tests {
|
||||
use std::time::Duration;
|
||||
use tokio::sync::watch;
|
||||
|
||||
#[test]
|
||||
fn rollout_probe_marker_is_only_trusted_from_a_forwarding_gateway() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
crate::tunnel::TUNNEL_RELAY_ROLLOUT_PROBE_HEADER,
|
||||
crate::tunnel::TUNNEL_RELAY_ROLLOUT_PROBE_VALUE
|
||||
.parse()
|
||||
.expect("probe marker should be a valid header"),
|
||||
);
|
||||
|
||||
assert!(!is_rollout_probe_request(&headers, false));
|
||||
assert!(is_rollout_probe_request(&headers, true));
|
||||
}
|
||||
|
||||
fn test_app_state() -> AppState {
|
||||
AppState::new(
|
||||
ControlPlaneClient::disabled(),
|
||||
|
||||
@@ -6,6 +6,9 @@ mod proxy_conn;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_gateway_tunnel::{
|
||||
resolve_proxy_max_streams, resolve_proxy_node_name, resolve_proxy_protocol_version,
|
||||
};
|
||||
use aether_runtime::{
|
||||
hold_admission_permit_until, prometheus_response, service_up_sample, AdmissionPermit,
|
||||
ConcurrencyError, ConcurrencyGate, ConcurrencySnapshot, MetricKind, MetricLabel, MetricSample,
|
||||
@@ -17,7 +20,6 @@ use axum::http::HeaderMap;
|
||||
use axum::response::{IntoResponse, Json};
|
||||
use axum::routing::{get, post};
|
||||
use axum::Router;
|
||||
use base64::Engine as _;
|
||||
use dashmap::DashMap;
|
||||
use tracing::warn;
|
||||
|
||||
@@ -335,121 +337,3 @@ pub async fn ws_proxy(
|
||||
})
|
||||
.into_response()
|
||||
}
|
||||
|
||||
fn resolve_proxy_max_streams(headers: &HeaderMap, fallback: usize) -> usize {
|
||||
headers
|
||||
.get("x-tunnel-max-streams")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| value.parse::<usize>().ok())
|
||||
.unwrap_or(fallback)
|
||||
.clamp(1, 2048)
|
||||
}
|
||||
|
||||
fn resolve_proxy_node_name(headers: &HeaderMap, node_id: &str) -> String {
|
||||
if let Some(decoded) = headers
|
||||
.get(aether_contracts::tunnel::TUNNEL_NODE_NAME_B64_HEADER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| {
|
||||
base64::engine::general_purpose::URL_SAFE_NO_PAD
|
||||
.decode(value.trim())
|
||||
.ok()
|
||||
})
|
||||
.and_then(|bytes| String::from_utf8(bytes).ok())
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty() && value.chars().count() <= 100)
|
||||
{
|
||||
return decoded;
|
||||
}
|
||||
|
||||
headers
|
||||
.get("x-node-name")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or(node_id)
|
||||
.to_string()
|
||||
}
|
||||
|
||||
fn resolve_proxy_protocol_version(headers: &HeaderMap) -> u8 {
|
||||
headers
|
||||
.get(aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| value.parse::<u8>().ok())
|
||||
.filter(|value| *value >= 1)
|
||||
.map(|value| value.min(aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION))
|
||||
.unwrap_or(1)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use axum::http::{HeaderMap, HeaderValue};
|
||||
use base64::Engine as _;
|
||||
|
||||
use super::{
|
||||
resolve_proxy_max_streams, resolve_proxy_node_name, resolve_proxy_protocol_version,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn proxy_max_streams_honors_small_advertised_capacity() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-tunnel-max-streams", HeaderValue::from_static("8"));
|
||||
|
||||
assert_eq!(resolve_proxy_max_streams(&headers, 128), 8);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn proxy_max_streams_caps_unreasonably_large_capacity() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-tunnel-max-streams", HeaderValue::from_static("9999"));
|
||||
|
||||
assert_eq!(resolve_proxy_max_streams(&headers, 128), 2048);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn proxy_protocol_version_defaults_to_v1_when_header_missing() {
|
||||
let headers = HeaderMap::new();
|
||||
assert_eq!(resolve_proxy_protocol_version(&headers), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn proxy_protocol_version_reads_advertised_version() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER,
|
||||
HeaderValue::from_static("2"),
|
||||
);
|
||||
|
||||
assert_eq!(resolve_proxy_protocol_version(&headers), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn proxy_node_name_reads_legacy_ascii_header() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-node-name", HeaderValue::from_static("edge-1"));
|
||||
|
||||
assert_eq!(resolve_proxy_node_name(&headers, "node-1"), "edge-1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn proxy_node_name_decodes_base64_header() {
|
||||
let mut headers = HeaderMap::new();
|
||||
let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode("日本节点");
|
||||
headers.insert(
|
||||
aether_contracts::tunnel::TUNNEL_NODE_NAME_B64_HEADER,
|
||||
HeaderValue::from_str(&encoded).expect("encoded header value should parse"),
|
||||
);
|
||||
|
||||
assert_eq!(resolve_proxy_node_name(&headers, "node-1"), "日本节点");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn proxy_node_name_falls_back_to_node_id_for_invalid_base64() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
aether_contracts::tunnel::TUNNEL_NODE_NAME_B64_HEADER,
|
||||
HeaderValue::from_static("not valid"),
|
||||
);
|
||||
|
||||
assert_eq!(resolve_proxy_node_name(&headers, "node-1"), "node-1");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,23 +1 @@
|
||||
use bytes::Bytes;
|
||||
|
||||
pub use aether_contracts::tunnel::{
|
||||
decode_payload, encode_connection_close, encode_frame, encode_goaway, encode_goaway_v3,
|
||||
encode_hello, encode_load_report, encode_ping, encode_pong, encode_reset_stream,
|
||||
encode_settings, encode_stream_error, encode_window_update, frame_payload_by_header,
|
||||
ConnectionClosePayload, FrameHeader, GoAwayPayload, HelloPayload, LoadReportPayload,
|
||||
RequestMeta, ResetStreamPayload, ResponseMeta, SettingsPayload, WindowUpdatePayload,
|
||||
CONNECTION_CLOSE, FLAG_END_STREAM, FLAG_GZIP_COMPRESSED, GOAWAY, HEADER_SIZE, HEARTBEAT_ACK,
|
||||
HEARTBEAT_DATA, HELLO, LOAD_REPORT, PING, PONG, REQUEST_BODY, REQUEST_HEADERS, RESET_STREAM,
|
||||
RESPONSE_BODY, RESPONSE_HEADERS, SETTINGS, STREAM_END, STREAM_ERROR, WINDOW_UPDATE,
|
||||
};
|
||||
|
||||
pub fn compress_payload(payload: &[u8]) -> Result<(Vec<u8>, u8), std::io::Error> {
|
||||
let (compressed, flags) =
|
||||
aether_contracts::tunnel::compress_payload(Bytes::copy_from_slice(payload));
|
||||
Ok((compressed.to_vec(), flags))
|
||||
}
|
||||
|
||||
pub fn raw_payload(payload: &[u8]) -> (Vec<u8>, u8) {
|
||||
let (payload, flags) = aether_contracts::tunnel::raw_payload(Bytes::copy_from_slice(payload));
|
||||
(payload.to_vec(), flags)
|
||||
}
|
||||
pub use aether_gateway_tunnel::embedded::protocol::*;
|
||||
|
||||
Reference in New Issue
Block a user