feat(security): harden gateway boundaries and usage policies

Consolidate subscription usage policy enforcement, privacy-safe persistence, and gateway security hardening into one reviewable change.

Includes bounded HTTP and execution envelopes, header and protocol guards, DNS and relay validation, authentication and secret projection hardening, secure backup/install paths, and regression coverage.
This commit is contained in:
elky
2026-09-04 03:45:52 +08:00
parent ddcbeb3ae9
commit 579f2c7cc1
1019 changed files with 190437 additions and 26080 deletions
@@ -0,0 +1,141 @@
use base64::Engine as _;
use hmac::{Hmac, Mac};
use sha2::{Digest as _, Sha256};
pub const INTERNAL_GATEWAY_AUTH_TIMESTAMP_HEADER: &str = "x-aether-internal-gateway-timestamp";
pub const INTERNAL_GATEWAY_AUTH_NONCE_HEADER: &str = "x-aether-internal-gateway-nonce";
pub const INTERNAL_GATEWAY_AUTH_SIGNATURE_HEADER: &str = "x-aether-internal-gateway-signature";
const INTERNAL_GATEWAY_AUTH_CONTEXT: &[u8] = b"aether-internal-gateway-auth-v1";
type HmacSha256 = Hmac<Sha256>;
pub fn sign_internal_gateway_request(
secret: &[u8],
method: &str,
path_and_query: &str,
timestamp_unix_secs: u64,
nonce: &str,
body: &[u8],
) -> String {
let mut mac = HmacSha256::new_from_slice(secret).expect("HMAC accepts keys of any size");
update_auth_mac(
&mut mac,
method,
path_and_query,
timestamp_unix_secs,
nonce,
body,
);
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes())
}
pub fn verify_internal_gateway_request_signature(
secret: &[u8],
method: &str,
path_and_query: &str,
timestamp_unix_secs: u64,
nonce: &str,
body: &[u8],
signature: &str,
) -> bool {
let signature = signature.trim();
if signature.len() > 43 {
return false;
}
let Ok(signature) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(signature) else {
return false;
};
let Ok(mut mac) = HmacSha256::new_from_slice(secret) else {
return false;
};
update_auth_mac(
&mut mac,
method,
path_and_query,
timestamp_unix_secs,
nonce,
body,
);
mac.verify_slice(&signature).is_ok()
}
fn update_auth_mac(
mac: &mut HmacSha256,
method: &str,
path_and_query: &str,
timestamp_unix_secs: u64,
nonce: &str,
body: &[u8],
) {
mac.update(INTERNAL_GATEWAY_AUTH_CONTEXT);
update_auth_field(mac, method.as_bytes());
update_auth_field(mac, path_and_query.as_bytes());
mac.update(&timestamp_unix_secs.to_be_bytes());
update_auth_field(mac, nonce.as_bytes());
mac.update(&(body.len() as u64).to_be_bytes());
mac.update(&Sha256::digest(body));
}
fn update_auth_field(mac: &mut HmacSha256, value: &[u8]) {
mac.update(&(value.len() as u64).to_be_bytes());
mac.update(value);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn signature_binds_every_security_relevant_field() {
let secret = b"internal-gateway-test-secret-32-bytes-minimum";
let body = br#"{"path":"/v1/models"}"#;
let timestamp = 1_800_000_000;
let nonce = "nonce-value-0000000000000001";
let path = "/api/internal/gateway/resolve?mode=full";
let signature = sign_internal_gateway_request(secret, "POST", path, timestamp, nonce, body);
assert!(verify_internal_gateway_request_signature(
secret, "POST", path, timestamp, nonce, body, &signature,
));
assert!(!verify_internal_gateway_request_signature(
secret, "GET", path, timestamp, nonce, body, &signature,
));
assert!(!verify_internal_gateway_request_signature(
secret,
"POST",
"/api/internal/gateway/resolve?mode=brief",
timestamp,
nonce,
body,
&signature,
));
assert!(!verify_internal_gateway_request_signature(
secret,
"POST",
path,
timestamp + 1,
nonce,
body,
&signature,
));
assert!(!verify_internal_gateway_request_signature(
secret,
"POST",
path,
timestamp,
"different-nonce-0000000000001",
body,
&signature,
));
assert!(!verify_internal_gateway_request_signature(
secret,
"POST",
path,
timestamp,
nonce,
br#"{"path":"/v1/providers"}"#,
&signature,
));
}
}
+6 -4
View File
@@ -1,5 +1,6 @@
mod error;
mod frame;
pub mod internal_gateway;
mod plan;
mod result;
pub mod tunnel;
@@ -9,13 +10,14 @@ mod usage;
pub use error::{ExecutionError, ExecutionErrorKind, ExecutionPhase};
pub use frame::{StreamFrame, StreamFramePayload, StreamFrameType};
pub use plan::{
ExecutionPlan, ExecutionResponseBodyMode, ExecutionTimeouts, ProxySnapshot, RequestBody,
ResolvedTransportProfile, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
redact_url_for_debug, ExecutionPlan, ExecutionResponseBodyMode, ExecutionTimeouts,
ProxySnapshot, RequestBody, ResolvedTransportProfile,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
EXECUTION_RESPONSE_BODY_MODE_HEADER, MAX_EXECUTION_REQUEST_TIMEOUT_MS,
MAX_EXECUTION_REQUEST_TIMEOUT_SECS, MAX_EXECUTION_STREAM_FIRST_BYTE_TIMEOUT_MS,
MAX_EXECUTION_STREAM_FIRST_BYTE_TIMEOUT_SECS, TRANSPORT_BACKEND_BROWSER_WREQ,
TRANSPORT_BACKEND_HYPER_RUSTLS, TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_AUTO,
MAX_EXECUTION_STREAM_FIRST_BYTE_TIMEOUT_SECS, PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY,
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_HYPER_RUSTLS,
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_AUTO,
TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
TRANSPORT_POOL_SCOPE_KEY,
};
+186 -6
View File
@@ -1,12 +1,11 @@
use std::collections::BTreeMap;
use std::fmt;
use serde::{Deserialize, Serialize};
use serde_json::Value;
pub const EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER: &str = "x-aether-execution-follow-redirects";
pub const EXECUTION_REQUEST_HTTP1_ONLY_HEADER: &str = "x-aether-execution-http1-only";
pub const EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER: &str =
"x-aether-execution-accept-invalid-certs";
pub const EXECUTION_RESPONSE_BODY_MODE_HEADER: &str = "x-aether-execution-response-body-mode";
pub const MAX_EXECUTION_REQUEST_TIMEOUT_SECS: u64 = 1_200;
pub const MAX_EXECUTION_REQUEST_TIMEOUT_MS: u64 = MAX_EXECUTION_REQUEST_TIMEOUT_SECS * 1_000;
@@ -57,7 +56,7 @@ pub struct ExecutionTimeouts {
pub total_ms: Option<u64>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[derive(Clone, Serialize, Deserialize, PartialEq)]
pub struct RequestBody {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub json_body: Option<Value>,
@@ -67,6 +66,27 @@ pub struct RequestBody {
pub body_ref: Option<String>,
}
impl fmt::Debug for RequestBody {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("RequestBody")
.field("has_json_body", &self.json_body.is_some())
.field(
"json_body_bytes",
&self
.json_body
.as_ref()
.and_then(|body| serde_json::to_vec(body).ok().map(|bytes| bytes.len())),
)
.field(
"body_bytes_b64_len",
&self.body_bytes_b64.as_ref().map(String::len),
)
.field("body_ref_len", &self.body_ref.as_ref().map(String::len))
.finish()
}
}
impl RequestBody {
pub fn from_json(json_body: Value) -> Self {
Self {
@@ -77,7 +97,7 @@ impl RequestBody {
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
#[derive(Clone, Default, Serialize, Deserialize, PartialEq)]
pub struct ProxySnapshot {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub enabled: Option<bool>,
@@ -93,6 +113,24 @@ pub struct ProxySnapshot {
pub extra: Option<Value>,
}
impl fmt::Debug for ProxySnapshot {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ProxySnapshot")
.field("enabled", &self.enabled)
.field("mode", &self.mode)
.field("node_id", &self.node_id)
.field("label", &self.label)
.field("url", &self.url.as_deref().map(redact_url_for_debug))
.field("has_extra", &self.extra.is_some())
.finish()
}
}
/// Reserved internal metadata key used to fence proxy-node mutations against
/// a node incarnation recreated under the same stable id.
pub const PROXY_NODE_TUNNEL_GENERATION_EXTRA_KEY: &str = "proxy_node_tunnel_generation";
pub const TRANSPORT_BACKEND_REQWEST_RUSTLS: &str = "reqwest_rustls";
pub const TRANSPORT_BACKEND_HYPER_RUSTLS: &str = "hyper_rustls";
pub const TRANSPORT_BACKEND_BROWSER_WREQ: &str = "browser_wreq";
@@ -101,7 +139,7 @@ pub const TRANSPORT_HTTP_MODE_HTTP1_ONLY: &str = "http1_only";
pub const TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE: &str = "h2c_prior_knowledge";
pub const TRANSPORT_POOL_SCOPE_KEY: &str = "key";
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[derive(Clone, Serialize, Deserialize, PartialEq)]
#[serde(default)]
pub struct ResolvedTransportProfile {
pub profile_id: String,
@@ -127,7 +165,21 @@ impl Default for ResolvedTransportProfile {
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
impl fmt::Debug for ResolvedTransportProfile {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ResolvedTransportProfile")
.field("profile_id", &self.profile_id)
.field("backend", &self.backend)
.field("http_mode", &self.http_mode)
.field("pool_scope", &self.pool_scope)
.field("has_header_fingerprint", &self.header_fingerprint.is_some())
.field("has_extra", &self.extra.is_some())
.finish()
}
}
#[derive(Clone, Serialize, Deserialize, PartialEq)]
pub struct ExecutionPlan {
pub request_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
@@ -162,6 +214,63 @@ pub struct ExecutionPlan {
pub timeouts: Option<ExecutionTimeouts>,
}
impl fmt::Debug for ExecutionPlan {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ExecutionPlan")
.field("request_id", &self.request_id)
.field("candidate_id", &self.candidate_id)
.field("provider_name", &self.provider_name)
.field("provider_id", &self.provider_id)
.field("endpoint_id", &self.endpoint_id)
.field("key_id", &self.key_id)
.field("method", &self.method)
.field("url", &redact_url_for_debug(&self.url))
.field("header_names", &self.headers.keys().collect::<Vec<_>>())
.field("content_type", &self.content_type)
.field("content_encoding", &self.content_encoding)
.field("body", &self.body)
.field("stream", &self.stream)
.field("client_api_format", &self.client_api_format)
.field("provider_api_format", &self.provider_api_format)
.field("model_name", &self.model_name)
.field("proxy", &self.proxy)
.field("transport_profile", &self.transport_profile)
.field("timeouts", &self.timeouts)
.finish()
}
}
/// Return a bounded URL representation suitable for diagnostics.
///
/// URL userinfo, query parameters, and fragments are never emitted because
/// providers commonly put API keys or OAuth tokens in those locations. An
/// unparsable URL is represented only by its length rather than echoing input.
pub fn redact_url_for_debug(raw: &str) -> String {
const MAX_DEBUG_URL_CHARS: usize = 512;
let raw = raw.trim();
if raw.is_empty() {
return String::new();
}
let Ok(mut url) = url::Url::parse(raw) else {
return format!("[invalid-url len={}]", raw.len());
};
let _ = url.set_username("");
let _ = url.set_password(None);
url.set_query(None);
url.set_fragment(None);
let rendered = url.to_string();
if rendered.chars().count() <= MAX_DEBUG_URL_CHARS {
rendered
} else {
let prefix = rendered
.chars()
.take(MAX_DEBUG_URL_CHARS - 3)
.collect::<String>();
format!("{prefix}...")
}
}
#[cfg(test)]
mod tests {
use super::*;
@@ -268,4 +377,75 @@ mod tests {
Some(300_000)
);
}
#[test]
fn debug_redacts_plan_credentials_and_payloads() {
let plan = ExecutionPlan {
request_id: "req-debug".into(),
candidate_id: Some("candidate-debug".into()),
provider_name: Some("provider".into()),
provider_id: "provider-id".into(),
endpoint_id: "endpoint-id".into(),
key_id: "key-id".into(),
method: "POST".into(),
url: "https://proxy-user:[email protected]/v1?api_key=url-secret#fragment-secret".into(),
headers: BTreeMap::from([(
"authorization".into(),
"Bearer header-secret".into(),
)]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(serde_json::json!({
"access_token": "body-secret",
"prompt": "hello"
})),
stream: false,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
model_name: Some("model".into()),
proxy: Some(ProxySnapshot {
enabled: Some(true),
mode: Some("http".into()),
node_id: Some("node".into()),
label: None,
url: Some("http://proxy-user:[email protected]?token=proxy-secret".into()),
extra: Some(serde_json::json!({"credential": "extra-secret"})),
}),
transport_profile: Some(ResolvedTransportProfile {
profile_id: "profile".into(),
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.into(),
http_mode: TRANSPORT_HTTP_MODE_AUTO.into(),
pool_scope: TRANSPORT_POOL_SCOPE_KEY.into(),
header_fingerprint: Some(serde_json::json!({"authorization": "fingerprint-secret"})),
extra: Some(serde_json::json!({"secret": "profile-secret"})),
}),
timeouts: None,
};
let debug = format!("{plan:?}");
for secret in [
"proxy-user",
"proxy-password",
"url-secret",
"fragment-secret",
"header-secret",
"body-secret",
"proxy-secret",
"extra-secret",
"fingerprint-secret",
"profile-secret",
] {
assert!(!debug.contains(secret), "debug leaked {secret}: {debug}");
}
assert!(debug.contains("header_names"));
assert!(debug.contains("has_json_body"));
assert!(debug.contains("profile"));
}
#[test]
fn debug_url_redaction_fails_closed_for_invalid_urls() {
let redacted = redact_url_for_debug("not a url?token=secret");
assert!(!redacted.contains("secret"));
assert!(redacted.starts_with("[invalid-url"));
}
}
+93 -2
View File
@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::fmt;
use serde::{Deserialize, Serialize};
use serde_json::Value;
@@ -22,7 +23,7 @@ pub struct ExecutionTelemetry {
pub upstream_bytes: Option<u64>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[derive(Clone, Serialize, Deserialize, PartialEq)]
pub struct ResponseBody {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub json_body: Option<Value>,
@@ -30,7 +31,27 @@ pub struct ResponseBody {
pub body_bytes_b64: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
impl fmt::Debug for ResponseBody {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ResponseBody")
.field("has_json_body", &self.json_body.is_some())
.field(
"json_body_bytes",
&self
.json_body
.as_ref()
.and_then(|body| serde_json::to_vec(body).ok().map(|bytes| bytes.len())),
)
.field(
"body_bytes_b64_len",
&self.body_bytes_b64.as_ref().map(String::len),
)
.finish()
}
}
#[derive(Clone, Serialize, Deserialize, PartialEq)]
pub struct ExecutionResult {
pub request_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
@@ -47,3 +68,73 @@ pub struct ExecutionResult {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error: Option<ExecutionError>,
}
impl fmt::Debug for ExecutionResult {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut debug = formatter.debug_struct("ExecutionResult");
debug
.field("request_id", &self.request_id)
.field("candidate_id", &self.candidate_id)
.field("status_code", &self.status_code)
.field("header_names", &self.headers.keys().collect::<Vec<_>>())
.field("response_observation", &self.response_observation)
.field("body", &self.body)
.field("telemetry", &self.telemetry)
.field("has_error", &self.error.is_some());
if let Some(error) = self.error.as_ref() {
// ExecutionError::message can contain an upstream response or URL.
debug
.field("error_kind", &error.kind)
.field("error_phase", &error.phase)
.field("error_upstream_status", &error.upstream_status)
.field("error_retryable", &error.retryable)
.field("error_failover_recommended", &error.failover_recommended);
}
debug.finish()
}
}
#[cfg(test)]
mod tests {
use super::{ExecutionResult, ResponseBody};
use crate::{ExecutionError, ExecutionErrorKind, ExecutionPhase};
use std::collections::BTreeMap;
#[test]
fn debug_does_not_render_response_headers_body_or_error_message() {
let result = ExecutionResult {
request_id: "request-1".into(),
candidate_id: None,
status_code: 401,
headers: BTreeMap::from([(
"set-cookie".into(),
"session=response-header-secret".into(),
)]),
response_observation: None,
body: Some(ResponseBody {
json_body: Some(serde_json::json!({"access_token": "response-body-secret"})),
body_bytes_b64: None,
}),
telemetry: None,
error: Some(ExecutionError {
kind: ExecutionErrorKind::Upstream4xx,
phase: ExecutionPhase::Finalize,
message: "upstream detail error-secret".into(),
upstream_status: Some(401),
retryable: false,
failover_recommended: false,
}),
};
let debug = format!("{result:?}");
for secret in [
"response-header-secret",
"response-body-secret",
"error-secret",
] {
assert!(!debug.contains(secret), "debug leaked {secret}: {debug}");
}
assert!(debug.contains("header_names"));
assert!(debug.contains("has_error"));
}
}
+500 -21
View File
@@ -1,18 +1,195 @@
use std::fmt;
use std::io::Read;
use base64::Engine as _;
use bytes::{Buf, BufMut, Bytes, BytesMut};
use flate2::read::GzDecoder;
use flate2::write::GzEncoder;
use flate2::Compression;
use hmac::{Hmac, Mac};
use sha2::{Digest as _, Sha256};
pub const HEADER_SIZE: usize = 10;
pub const TUNNEL_RELAY_FORWARDED_BY_HEADER: &str = "x-aether-tunnel-forwarded-by";
pub const TUNNEL_RELAY_OWNER_INSTANCE_HEADER: &str = "x-aether-tunnel-owner-instance-id";
pub const TUNNEL_RELAY_AUTH_SENDER_HEADER: &str = "x-aether-tunnel-relay-sender";
pub const TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER: &str = "x-aether-tunnel-relay-timestamp";
pub const TUNNEL_RELAY_AUTH_NONCE_HEADER: &str = "x-aether-tunnel-relay-nonce";
pub const TUNNEL_RELAY_AUTH_PAYLOAD_HEADER: &str = "x-aether-tunnel-relay-payload";
pub const TUNNEL_RELAY_AUTH_SIGNATURE_HEADER: &str = "x-aether-tunnel-relay-signature";
pub const TUNNEL_PROTOCOL_VERSION_HEADER: &str = "x-aether-tunnel-protocol-version";
pub const TUNNEL_NODE_NAME_B64_HEADER: &str = "x-aether-tunnel-node-name-b64";
pub const CURRENT_TUNNEL_PROTOCOL_VERSION: u8 = 3;
pub const CURRENT_TUNNEL_PROTOCOL_VERSION_STR: &str = "3";
pub const MAX_TUNNEL_RELAY_META_LEN: usize = 256 * 1024;
/// Keep decoded tunnel frames within the same size envelope enforced by the
/// WebSocket transports. This also bounds gzip expansion for untrusted peers.
pub const MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES: usize = 64 * 1024 * 1024;
const TUNNEL_RELAY_AUTH_CONTEXT: &[u8] = b"aether-tunnel-relay-auth-v2";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TunnelRelayPayloadDigest {
metadata_sha256: [u8; 32],
body_len: u64,
body_sha256: [u8; 32],
}
impl TunnelRelayPayloadDigest {
pub fn body_len(self) -> u64 {
self.body_len
}
pub fn encode_header_value(self) -> String {
let mut encoded = [0_u8; 72];
encoded[..32].copy_from_slice(&self.metadata_sha256);
encoded[32..40].copy_from_slice(&self.body_len.to_be_bytes());
encoded[40..].copy_from_slice(&self.body_sha256);
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(encoded)
}
pub fn decode_header_value(value: &str) -> Option<Self> {
let value = value.trim();
if value.len() > 96 {
return None;
}
let decoded = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(value)
.ok()?;
let encoded: [u8; 72] = decoded.try_into().ok()?;
Some(Self {
metadata_sha256: encoded[..32].try_into().ok()?,
body_len: u64::from_be_bytes(encoded[32..40].try_into().ok()?),
body_sha256: encoded[40..].try_into().ok()?,
})
}
pub fn matches_metadata(self, metadata_envelope: &[u8]) -> bool {
self.metadata_sha256 == <[u8; 32]>::from(Sha256::digest(metadata_envelope))
}
pub fn matches_body(self, body: &[u8]) -> bool {
self.matches_body_hash(body.len() as u64, Sha256::digest(body).into())
}
pub fn matches_body_hash(self, body_len: u64, body_sha256: [u8; 32]) -> bool {
self.body_len == body_len && self.body_sha256 == body_sha256
}
}
pub fn tunnel_relay_payload_digest(
metadata_envelope: &[u8],
body: &[u8],
) -> TunnelRelayPayloadDigest {
TunnelRelayPayloadDigest {
metadata_sha256: Sha256::digest(metadata_envelope).into(),
body_len: body.len() as u64,
body_sha256: Sha256::digest(body).into(),
}
}
pub fn tunnel_relay_payload_digest_from_hashes(
metadata_envelope: &[u8],
body_len: u64,
body_sha256: [u8; 32],
) -> TunnelRelayPayloadDigest {
TunnelRelayPayloadDigest {
metadata_sha256: Sha256::digest(metadata_envelope).into(),
body_len,
body_sha256,
}
}
pub fn sign_tunnel_relay_request(
secret: &[u8],
sender_instance_id: &str,
owner_instance_id: &str,
node_id: &str,
forwarded_by: &str,
rollout_probe: bool,
timestamp_unix_secs: u64,
nonce: &str,
payload_digest: &TunnelRelayPayloadDigest,
) -> String {
let mut mac = Hmac::<Sha256>::new_from_slice(secret).expect("HMAC accepts keys of any size");
update_tunnel_relay_auth_mac(
&mut mac,
sender_instance_id,
owner_instance_id,
node_id,
forwarded_by,
rollout_probe,
timestamp_unix_secs,
nonce,
payload_digest,
);
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes())
}
pub fn verify_tunnel_relay_request_signature(
secret: &[u8],
sender_instance_id: &str,
owner_instance_id: &str,
node_id: &str,
forwarded_by: &str,
rollout_probe: bool,
timestamp_unix_secs: u64,
nonce: &str,
payload_digest: &TunnelRelayPayloadDigest,
signature: &str,
) -> bool {
let signature = signature.trim();
if signature.len() > 43 {
return false;
}
let Ok(signature) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(signature) else {
return false;
};
let Ok(mut mac) = Hmac::<Sha256>::new_from_slice(secret) else {
return false;
};
update_tunnel_relay_auth_mac(
&mut mac,
sender_instance_id,
owner_instance_id,
node_id,
forwarded_by,
rollout_probe,
timestamp_unix_secs,
nonce,
payload_digest,
);
mac.verify_slice(&signature).is_ok()
}
fn update_tunnel_relay_auth_mac(
mac: &mut Hmac<Sha256>,
sender_instance_id: &str,
owner_instance_id: &str,
node_id: &str,
forwarded_by: &str,
rollout_probe: bool,
timestamp_unix_secs: u64,
nonce: &str,
payload_digest: &TunnelRelayPayloadDigest,
) {
mac.update(TUNNEL_RELAY_AUTH_CONTEXT);
update_tunnel_relay_auth_field(mac, sender_instance_id.as_bytes());
update_tunnel_relay_auth_field(mac, owner_instance_id.as_bytes());
update_tunnel_relay_auth_field(mac, node_id.as_bytes());
update_tunnel_relay_auth_field(mac, forwarded_by.as_bytes());
mac.update(&[u8::from(rollout_probe)]);
mac.update(&timestamp_unix_secs.to_be_bytes());
update_tunnel_relay_auth_field(mac, nonce.as_bytes());
mac.update(&payload_digest.metadata_sha256);
mac.update(&payload_digest.body_len.to_be_bytes());
mac.update(&payload_digest.body_sha256);
}
fn update_tunnel_relay_auth_field(mac: &mut Hmac<Sha256>, value: &[u8]) {
mac.update(&(value.len() as u64).to_be_bytes());
mac.update(value);
}
pub mod flags {
pub const END_STREAM: u8 = 0x01;
@@ -111,7 +288,7 @@ impl FrameHeader {
}
}
#[derive(Debug, Clone)]
#[derive(Clone)]
pub struct Frame {
pub stream_id: u32,
pub msg_type: MsgType,
@@ -119,6 +296,18 @@ pub struct Frame {
pub payload: Bytes,
}
impl fmt::Debug for Frame {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("Frame")
.field("stream_id", &self.stream_id)
.field("msg_type", &self.msg_type)
.field("flags", &self.flags)
.field("payload_len", &self.payload.len())
.finish()
}
}
impl Frame {
pub fn new(stream_id: u32, msg_type: MsgType, flags: u8, payload: impl Into<Bytes>) -> Self {
Self {
@@ -170,6 +359,12 @@ impl Frame {
actual: HEADER_SIZE + data.remaining(),
});
}
if data.remaining() > payload_len {
return Err(ProtocolError::Trailing {
expected: HEADER_SIZE + payload_len,
actual: HEADER_SIZE + data.remaining(),
});
}
let msg_type =
MsgType::from_u8(msg_type_raw).ok_or(ProtocolError::UnknownMsgType(msg_type_raw))?;
@@ -190,11 +385,13 @@ pub enum ProtocolError {
TooShort { expected: usize, actual: usize },
#[error("frame incomplete: expected {expected} bytes, got {actual}")]
Incomplete { expected: usize, actual: usize },
#[error("frame has trailing bytes: expected {expected} bytes, got {actual}")]
Trailing { expected: usize, actual: usize },
#[error("unknown message type: 0x{0:02x}")]
UnknownMsgType(u8),
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
#[derive(Clone, serde::Serialize, serde::Deserialize)]
pub struct RequestMeta {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_id: Option<String>,
@@ -221,6 +418,30 @@ pub struct RequestMeta {
pub transport_profile: Option<crate::ResolvedTransportProfile>,
}
impl fmt::Debug for RequestMeta {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("RequestMeta")
.field("provider_id", &self.provider_id)
.field("endpoint_id", &self.endpoint_id)
.field("key_id", &self.key_id)
.field("method", &self.method)
.field("url", &crate::redact_url_for_debug(&self.url))
.field("header_names", &self.headers.keys().collect::<Vec<_>>())
.field("stream", &self.stream)
.field("request_timeout_ms", &self.request_timeout_ms)
.field(
"stream_first_byte_timeout_ms",
&self.stream_first_byte_timeout_ms,
)
.field("timeout", &self.timeout)
.field("follow_redirects", &self.follow_redirects)
.field("http1_only", &self.http1_only)
.field("transport_profile", &self.transport_profile)
.finish()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ResolvedTunnelRequestTimeouts {
pub first_byte_ms: u64,
@@ -316,12 +537,29 @@ where
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
#[derive(Clone, serde::Serialize, serde::Deserialize)]
pub struct ResponseMeta {
pub status: u16,
pub headers: Vec<(String, String)>,
}
impl fmt::Debug for ResponseMeta {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ResponseMeta")
.field("status", &self.status)
.field(
"header_names",
&self
.headers
.iter()
.map(|(name, _)| name)
.collect::<Vec<_>>(),
)
.finish()
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct HelloPayload {
pub protocol_version: u8,
@@ -453,30 +691,47 @@ fn encode_json_control<T: serde::Serialize>(msg_type: u8, payload: &T) -> Vec<u8
pub fn frame_payload_by_header<'a>(data: &'a [u8], header: &FrameHeader) -> Option<&'a [u8]> {
let payload_len = header.payload_len as usize;
let end = HEADER_SIZE.checked_add(payload_len)?;
if data.len() < end {
if data.len() != end {
return None;
}
Some(&data[HEADER_SIZE..end])
}
pub fn decode_payload(data: &[u8], header: &FrameHeader) -> Result<Vec<u8>, String> {
decode_payload_with_limit(data, header, MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES)
}
pub fn decode_payload_with_limit(
data: &[u8],
header: &FrameHeader,
max_decoded_bytes: usize,
) -> Result<Vec<u8>, String> {
let payload = frame_payload_by_header(data, header)
.ok_or_else(|| "incomplete frame payload".to_string())?;
if header.flags & FLAG_GZIP_COMPRESSED != 0 {
let mut decoder = GzDecoder::new(payload);
let mut decoded = Vec::new();
decoder
.read_to_end(&mut decoded)
.map_err(|err| format!("failed to decompress payload: {err}"))?;
Ok(decoded)
decompress_gzip_with_limit(payload, max_decoded_bytes)
.map_err(|err| format!("failed to decompress payload: {err}"))
} else if payload.len() > max_decoded_bytes {
Err(format!(
"decoded tunnel payload exceeds {max_decoded_bytes} bytes"
))
} else {
Ok(payload.to_vec())
}
}
pub fn decompress_if_gzip(frame: &Frame) -> Result<Bytes, std::io::Error> {
decompress_if_gzip_with_limit(frame, MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES)
}
pub fn decompress_if_gzip_with_limit(
frame: &Frame,
max_decoded_bytes: usize,
) -> Result<Bytes, std::io::Error> {
if frame.is_gzip() {
decompress_gzip(&frame.payload)
decompress_gzip_with_limit(&frame.payload, max_decoded_bytes).map(Bytes::from)
} else if frame.payload.len() > max_decoded_bytes {
Err(decoded_payload_too_large(max_decoded_bytes))
} else {
Ok(frame.payload.clone())
}
@@ -499,11 +754,43 @@ pub fn raw_payload(data: Bytes) -> (Bytes, u8) {
const COMPRESS_MIN_SIZE: usize = 512;
fn decompress_gzip(data: &[u8]) -> Result<Bytes, std::io::Error> {
fn decompress_gzip_with_limit(
data: &[u8],
max_decoded_bytes: usize,
) -> Result<Vec<u8>, std::io::Error> {
let mut decoder = GzDecoder::new(data);
let mut buf = Vec::new();
decoder.read_to_end(&mut buf)?;
Ok(Bytes::from(buf))
let mut decoded = Vec::with_capacity(max_decoded_bytes.min(8 * 1024));
let mut chunk = [0_u8; 8 * 1024];
loop {
let remaining = max_decoded_bytes.saturating_sub(decoded.len());
let read_len = remaining.saturating_add(1).min(chunk.len());
let read = match decoder.read(&mut chunk[..read_len]) {
Ok(read) => read,
Err(error) if error.kind() == std::io::ErrorKind::Interrupted => continue,
Err(error) => return Err(error),
};
if read == 0 {
return Ok(decoded);
}
if read > remaining {
return Err(decoded_payload_too_large(max_decoded_bytes));
}
if decoded.capacity().saturating_sub(decoded.len()) < read {
decoded.try_reserve_exact(read).map_err(|error| {
std::io::Error::other(format!(
"failed to allocate decoded tunnel payload: {error}"
))
})?;
}
decoded.extend_from_slice(&chunk[..read]);
}
}
fn decoded_payload_too_large(max_decoded_bytes: usize) -> std::io::Error {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("decoded tunnel payload exceeds {max_decoded_bytes} bytes"),
)
}
fn compress_gzip(data: &[u8]) -> Result<Bytes, std::io::Error> {
@@ -518,12 +805,15 @@ fn compress_gzip(data: &[u8]) -> Result<Bytes, std::io::Error> {
#[cfg(test)]
mod tests {
use super::{
compress_payload, decode_payload, encode_frame, encode_goaway_v3, encode_ping,
encode_reset_stream, encode_window_update, raw_payload, resolve_tunnel_request_timeouts,
try_decode_tunnel_relay_request_meta, Frame, FrameHeader, GoAwayPayload, MsgType,
RequestMeta, ResetStreamPayload, WindowUpdatePayload, CURRENT_TUNNEL_PROTOCOL_VERSION,
CURRENT_TUNNEL_PROTOCOL_VERSION_STR, FLAG_GZIP_COMPRESSED, MAX_TUNNEL_RELAY_META_LEN,
REQUEST_HEADERS, TUNNEL_PROTOCOL_VERSION_HEADER,
compress_payload, decode_payload, decode_payload_with_limit, decompress_if_gzip_with_limit,
encode_frame, encode_goaway_v3, encode_ping, encode_reset_stream, encode_window_update,
frame_payload_by_header, raw_payload, resolve_tunnel_request_timeouts,
sign_tunnel_relay_request, try_decode_tunnel_relay_request_meta,
tunnel_relay_payload_digest, verify_tunnel_relay_request_signature, Frame, FrameHeader,
GoAwayPayload, MsgType, ProtocolError, RequestMeta, ResetStreamPayload, ResponseMeta,
WindowUpdatePayload, CURRENT_TUNNEL_PROTOCOL_VERSION, CURRENT_TUNNEL_PROTOCOL_VERSION_STR,
FLAG_GZIP_COMPRESSED, HEADER_SIZE, MAX_TUNNEL_RELAY_META_LEN, REQUEST_HEADERS,
RESPONSE_BODY, TUNNEL_PROTOCOL_VERSION_HEADER,
};
use bytes::Bytes;
@@ -628,6 +918,71 @@ mod tests {
assert!(try_decode_tunnel_relay_request_meta(&oversized).is_err());
}
#[test]
fn tunnel_relay_signature_binds_routing_metadata_and_body() {
let digest = tunnel_relay_payload_digest(b"metadata", b"request-body");
let signature = sign_tunnel_relay_request(
b"shared-secret",
"gateway-a",
"gateway-b",
"node-1",
"gateway-a",
false,
123,
"nonce-1",
&digest,
);
assert!(verify_tunnel_relay_request_signature(
b"shared-secret",
"gateway-a",
"gateway-b",
"node-1",
"gateway-a",
false,
123,
"nonce-1",
&digest,
&signature,
));
assert!(!verify_tunnel_relay_request_signature(
b"shared-secret",
"gateway-a",
"gateway-b",
"node-2",
"gateway-a",
false,
123,
"nonce-1",
&digest,
&signature,
));
assert!(!verify_tunnel_relay_request_signature(
b"shared-secret",
"gateway-a",
"gateway-b",
"node-1",
"gateway-a",
false,
123,
"nonce-1",
&tunnel_relay_payload_digest(b"tampered", b"request-body"),
&signature,
));
assert!(!verify_tunnel_relay_request_signature(
b"shared-secret",
"gateway-a",
"gateway-b",
"node-1",
"gateway-a",
false,
123,
"nonce-1",
&tunnel_relay_payload_digest(b"metadata", b"tampered-body"),
&signature,
));
}
#[test]
fn request_meta_accepts_integer_timeout() {
let raw = br#"{"method":"GET","url":"https://example.com","headers":{},"timeout":15}"#;
@@ -652,6 +1007,34 @@ mod tests {
assert_eq!(decoded.payload, Bytes::from_static(b"hello"));
}
#[test]
fn frame_decode_rejects_trailing_bytes() {
let frame = Frame::new(7, MsgType::ResponseBody, 0, Bytes::from_static(b"hello"));
let mut encoded = frame.encode().to_vec();
encoded.extend_from_slice(b"hidden");
let error = Frame::decode(Bytes::from(encoded)).expect_err("trailing bytes rejected");
assert!(matches!(
error,
ProtocolError::Trailing { expected, actual }
if expected == HEADER_SIZE + 5 && actual == HEADER_SIZE + 11
));
}
#[test]
fn frame_payload_lookup_rejects_trailing_bytes() {
let encoded = encode_frame(7, RESPONSE_BODY, 0, b"hello");
let header = FrameHeader::parse(&encoded).expect("frame header should parse");
let mut with_trailing = encoded.clone();
with_trailing.push(0);
assert!(frame_payload_by_header(&with_trailing, &header).is_none());
assert_eq!(
frame_payload_by_header(&encoded, &header),
Some(&encoded[HEADER_SIZE..])
);
}
#[test]
fn frame_header_parses_raw_ping_frame() {
let encoded = encode_ping();
@@ -681,6 +1064,56 @@ mod tests {
assert_eq!(decoded, control_payload.to_vec());
}
#[test]
fn compressed_tunnel_payload_is_rejected_before_exceeding_decode_limit() {
const LIMIT: usize = 1024;
let at_limit = Bytes::from(vec![b'a'; LIMIT]);
let (at_limit_compressed, at_limit_flags) = compress_payload(at_limit.clone());
assert_ne!(at_limit_flags & FLAG_GZIP_COMPRESSED, 0);
let at_limit_frame = Frame::new(
1,
MsgType::RequestHeaders,
at_limit_flags,
at_limit_compressed,
);
assert_eq!(
decompress_if_gzip_with_limit(&at_limit_frame, LIMIT)
.expect("payload at the limit should decode"),
at_limit
);
let over_limit = Bytes::from(vec![b'a'; LIMIT + 1]);
let (compressed, flags) = compress_payload(over_limit);
assert_ne!(flags & FLAG_GZIP_COMPRESSED, 0);
let frame = Frame::new(1, MsgType::RequestHeaders, flags, compressed.clone());
let error = decompress_if_gzip_with_limit(&frame, LIMIT)
.expect_err("gzip expansion must be bounded");
assert_eq!(error.kind(), std::io::ErrorKind::InvalidData);
assert!(error.to_string().contains("exceeds 1024 bytes"));
let encoded = encode_frame(1, REQUEST_HEADERS, flags, &compressed);
let header = FrameHeader::parse(&encoded).expect("frame should parse");
let error = decode_payload_with_limit(&encoded, &header, LIMIT)
.expect_err("compatibility decoder must apply the same bound");
assert!(error.contains("exceeds 1024 bytes"));
let raw_at_limit = Frame::new(1, MsgType::RequestBody, 0, Bytes::from(vec![b'x'; LIMIT]));
assert!(decompress_if_gzip_with_limit(&raw_at_limit, LIMIT).is_ok());
let raw_over_limit = Frame::new(
1,
MsgType::RequestBody,
0,
Bytes::from(vec![b'x'; LIMIT + 1]),
);
assert_eq!(
decompress_if_gzip_with_limit(&raw_over_limit, LIMIT)
.expect_err("raw payloads must use the same bound")
.kind(),
std::io::ErrorKind::InvalidData
);
}
#[test]
fn tunnel_protocol_version_header_defaults_to_v2() {
assert_eq!(
@@ -719,4 +1152,50 @@ mod tests {
assert_eq!(goaway_payload.drain_deadline_ms, 30_000);
assert_eq!(goaway_payload.reason, "rolling restart");
}
#[test]
fn debug_does_not_render_tunnel_credentials_or_payload_bytes() {
let meta = RequestMeta {
provider_id: Some("provider".into()),
endpoint_id: Some("endpoint".into()),
key_id: Some("key".into()),
method: "POST".into(),
url: "https://user:[email protected]/path?token=url-secret".into(),
headers: std::collections::HashMap::from([(
"authorization".into(),
"Bearer header-secret".into(),
)]),
stream: false,
request_timeout_ms: None,
stream_first_byte_timeout_ms: None,
timeout: 60,
follow_redirects: None,
http1_only: false,
transport_profile: None,
};
let response = ResponseMeta {
status: 200,
headers: vec![("set-cookie".into(), "session=response-secret".into())],
};
let frame = Frame::new(
1,
MsgType::RequestBody,
0,
Bytes::from_static(b"request-body-secret"),
);
let debug = format!("{meta:?} {response:?} {frame:?}");
for secret in [
"user",
"password",
"url-secret",
"header-secret",
"response-secret",
"request-body-secret",
] {
assert!(!debug.contains(secret), "debug leaked {secret}: {debug}");
}
assert!(debug.contains("header_names"));
assert!(debug.contains("payload_len"));
}
}
+547 -4
View File
@@ -3,7 +3,7 @@ use aes_gcm::{Aes256Gcm, KeyInit, Nonce};
use base64::Engine;
use bytes::{Buf, BufMut, Bytes, BytesMut};
use hmac::{Hmac, Mac};
use sha2::Sha256;
use sha2::{Digest as _, Sha256};
use std::sync::atomic::{AtomicU64, Ordering};
use crate::tunnel::{Frame, MsgType, HEADER_SIZE};
@@ -12,10 +12,21 @@ type HmacSha256 = Hmac<Sha256>;
pub const TUNNEL_SECURITY_HEADER: &str = "x-aether-tunnel-security";
pub const TUNNEL_SECURITY_SESSION_HEADER: &str = "x-aether-tunnel-security-session";
pub const TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER: &str = "x-aether-tunnel-security-proof-timestamp";
pub const TUNNEL_SECURITY_PROOF_NONCE_HEADER: &str = "x-aether-tunnel-security-proof-nonce";
pub const TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER: &str = "x-aether-tunnel-security-proof-signature";
pub const TUNNEL_GENERATION_HEADER: &str = "x-aether-tunnel-generation";
pub const TUNNEL_CONTROL_PLANE_NODE_ID_HEADER: &str = "x-aether-tunnel-control-plane-node-id";
pub const TUNNEL_CONTROL_PLANE_GENERATION_HEADER: &str = "x-aether-tunnel-control-plane-generation";
pub const TUNNEL_CONTROL_PLANE_TIMESTAMP_HEADER: &str = "x-aether-tunnel-control-plane-timestamp";
pub const TUNNEL_CONTROL_PLANE_NONCE_HEADER: &str = "x-aether-tunnel-control-plane-nonce";
pub const TUNNEL_CONTROL_PLANE_SIGNATURE_HEADER: &str = "x-aether-tunnel-control-plane-signature";
pub const TUNNEL_SECURITY_NON_TLS_REQUIRED: &str = "non_tls_required";
pub const FLAG_ENCRYPTED: u8 = 0x04;
const CONTEXT: &[u8] = b"aether-tunnel-secure-v1";
const HANDSHAKE_PROOF_CONTEXT: &[u8] = b"aether-tunnel-handshake-proof-v1";
const CONTROL_PLANE_AUTH_CONTEXT: &[u8] = b"aether-tunnel-control-plane-auth-v1";
const CLIENT_TO_SERVER_LABEL: &[u8] = b"client-to-server";
const SERVER_TO_CLIENT_LABEL: &[u8] = b"server-to-client";
const CLIENT_TO_SERVER_NONCE_PREFIX: [u8; 4] = *b"c2s1";
@@ -41,6 +52,8 @@ pub enum TunnelSecurityError {
PayloadTooShort,
#[error("secure tunnel frame sequence is not the expected next value")]
UnexpectedSequence,
#[error("secure tunnel frame sequence space is exhausted")]
SequenceExhausted,
#[error("secure tunnel frame encryption failed")]
Encrypt,
#[error("secure tunnel frame decryption failed")]
@@ -98,7 +111,12 @@ impl SecureFrameCodec {
}
pub fn encrypt_frame(&self, frame: Frame) -> Result<Bytes, TunnelSecurityError> {
let sequence = self.next_sequence.fetch_add(1, Ordering::Relaxed);
let sequence = self
.next_sequence
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
current.checked_add(1)
})
.map_err(|_| TunnelSecurityError::SequenceExhausted)?;
let nonce_bytes = nonce_bytes(self.seal_prefix, sequence);
let nonce = Nonce::from_slice(&nonce_bytes);
let clear_flags = frame.flags & !FLAG_ENCRYPTED;
@@ -154,8 +172,17 @@ impl SecureFrameCodec {
},
)
.map_err(|_| TunnelSecurityError::Decrypt)?;
let next_sequence = expected_sequence
.checked_add(1)
.ok_or(TunnelSecurityError::SequenceExhausted)?;
self.next_open_sequence
.store(expected_sequence.wrapping_add(1), Ordering::Relaxed);
.compare_exchange(
expected_sequence,
next_sequence,
Ordering::AcqRel,
Ordering::Acquire,
)
.map_err(|_| TunnelSecurityError::UnexpectedSequence)?;
Ok(Frame::new(
frame.stream_id,
@@ -167,14 +194,275 @@ impl SecureFrameCodec {
}
pub fn decode_psk(key: &str) -> Result<[u8; 32], TunnelSecurityError> {
let key = key.trim();
if key.len() > 44 {
return Err(TunnelSecurityError::InvalidKey);
}
let decoded = base64::engine::general_purpose::STANDARD
.decode(key.trim())
.decode(key)
.map_err(|_| TunnelSecurityError::InvalidKey)?;
decoded
.try_into()
.map_err(|_| TunnelSecurityError::InvalidKey)
}
pub fn sign_tunnel_security_handshake(
key: &str,
node_id: &str,
security_mode: &str,
session_id: &str,
protocol_version: u8,
timestamp_unix_secs: u64,
nonce: &str,
) -> Result<String, TunnelSecurityError> {
sign_tunnel_security_handshake_for_generation(
key,
node_id,
"",
security_mode,
session_id,
protocol_version,
timestamp_unix_secs,
nonce,
)
}
pub fn sign_tunnel_security_handshake_for_generation(
key: &str,
node_id: &str,
tunnel_generation: &str,
security_mode: &str,
session_id: &str,
protocol_version: u8,
timestamp_unix_secs: u64,
nonce: &str,
) -> Result<String, TunnelSecurityError> {
let psk = decode_psk(key)?;
let mut mac =
<HmacSha256 as Mac>::new_from_slice(&psk).expect("HMAC accepts a 32-byte tunnel PSK");
update_handshake_proof_mac(
&mut mac,
node_id,
tunnel_generation,
security_mode,
session_id,
protocol_version,
timestamp_unix_secs,
nonce,
);
Ok(base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes()))
}
pub fn verify_tunnel_security_handshake(
key: &str,
node_id: &str,
security_mode: &str,
session_id: &str,
protocol_version: u8,
timestamp_unix_secs: u64,
nonce: &str,
signature: &str,
) -> bool {
verify_tunnel_security_handshake_for_generation(
key,
node_id,
"",
security_mode,
session_id,
protocol_version,
timestamp_unix_secs,
nonce,
signature,
)
}
pub fn verify_tunnel_security_handshake_for_generation(
key: &str,
node_id: &str,
tunnel_generation: &str,
security_mode: &str,
session_id: &str,
protocol_version: u8,
timestamp_unix_secs: u64,
nonce: &str,
signature: &str,
) -> bool {
let Ok(psk) = decode_psk(key) else {
return false;
};
if signature.len() > 43 {
return false;
}
let Ok(signature) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(signature) else {
return false;
};
let Ok(mut mac) = <HmacSha256 as Mac>::new_from_slice(&psk) else {
return false;
};
update_handshake_proof_mac(
&mut mac,
node_id,
tunnel_generation,
security_mode,
session_id,
protocol_version,
timestamp_unix_secs,
nonce,
);
mac.verify_slice(&signature).is_ok()
}
pub fn sign_tunnel_control_plane_request(
key: &str,
method: &str,
path: &str,
node_id: &str,
timestamp_unix_secs: u64,
nonce: &str,
body: &[u8],
) -> Result<String, TunnelSecurityError> {
sign_tunnel_control_plane_request_for_generation(
key,
method,
path,
node_id,
"",
timestamp_unix_secs,
nonce,
body,
)
}
pub fn sign_tunnel_control_plane_request_for_generation(
key: &str,
method: &str,
path: &str,
node_id: &str,
tunnel_generation: &str,
timestamp_unix_secs: u64,
nonce: &str,
body: &[u8],
) -> Result<String, TunnelSecurityError> {
let psk = decode_psk(key)?;
let mut mac =
<HmacSha256 as Mac>::new_from_slice(&psk).expect("HMAC accepts a 32-byte tunnel PSK");
update_control_plane_auth_mac(
&mut mac,
method,
path,
node_id,
tunnel_generation,
timestamp_unix_secs,
nonce,
body,
);
Ok(base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes()))
}
pub fn verify_tunnel_control_plane_request(
key: &str,
method: &str,
path: &str,
node_id: &str,
timestamp_unix_secs: u64,
nonce: &str,
body: &[u8],
signature: &str,
) -> bool {
verify_tunnel_control_plane_request_for_generation(
key,
method,
path,
node_id,
"",
timestamp_unix_secs,
nonce,
body,
signature,
)
}
pub fn verify_tunnel_control_plane_request_for_generation(
key: &str,
method: &str,
path: &str,
node_id: &str,
tunnel_generation: &str,
timestamp_unix_secs: u64,
nonce: &str,
body: &[u8],
signature: &str,
) -> bool {
let Ok(psk) = decode_psk(key) else {
return false;
};
if signature.len() > 43 {
return false;
}
let Ok(signature) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(signature) else {
return false;
};
let Ok(mut mac) = <HmacSha256 as Mac>::new_from_slice(&psk) else {
return false;
};
update_control_plane_auth_mac(
&mut mac,
method,
path,
node_id,
tunnel_generation,
timestamp_unix_secs,
nonce,
body,
);
mac.verify_slice(&signature).is_ok()
}
fn update_control_plane_auth_mac(
mac: &mut HmacSha256,
method: &str,
path: &str,
node_id: &str,
tunnel_generation: &str,
timestamp_unix_secs: u64,
nonce: &str,
body: &[u8],
) {
mac.update(CONTROL_PLANE_AUTH_CONTEXT);
update_handshake_proof_field(mac, method.as_bytes());
update_handshake_proof_field(mac, path.as_bytes());
update_handshake_proof_field(mac, node_id.as_bytes());
update_handshake_proof_field(mac, tunnel_generation.as_bytes());
mac.update(&timestamp_unix_secs.to_be_bytes());
update_handshake_proof_field(mac, nonce.as_bytes());
mac.update(&Sha256::digest(body));
}
fn update_handshake_proof_mac(
mac: &mut HmacSha256,
node_id: &str,
tunnel_generation: &str,
security_mode: &str,
session_id: &str,
protocol_version: u8,
timestamp_unix_secs: u64,
nonce: &str,
) {
mac.update(HANDSHAKE_PROOF_CONTEXT);
update_handshake_proof_field(mac, node_id.as_bytes());
update_handshake_proof_field(mac, tunnel_generation.as_bytes());
update_handshake_proof_field(mac, security_mode.as_bytes());
update_handshake_proof_field(mac, session_id.as_bytes());
mac.update(&[protocol_version]);
mac.update(&timestamp_unix_secs.to_be_bytes());
update_handshake_proof_field(mac, nonce.as_bytes());
}
fn update_handshake_proof_field(mac: &mut HmacSha256, value: &[u8]) {
mac.update(&(value.len() as u64).to_be_bytes());
mac.update(value);
}
fn derive_key(psk: &[u8; 32], session_id: &[u8], label: &[u8]) -> [u8; 32] {
let mut mac = <HmacSha256 as Mac>::new_from_slice(psk).expect("HMAC accepts 32-byte PSK");
mac.update(CONTEXT);
@@ -276,6 +564,70 @@ mod tests {
));
}
#[test]
fn secure_frame_accepts_a_concurrent_sequence_only_once() {
use std::sync::{Arc, Barrier};
let client = SecureFrameCodec::new(&test_key(), "session-1", TunnelSecurityRole::Client)
.expect("client codec");
let server = Arc::new(
SecureFrameCodec::new(&test_key(), "session-1", TunnelSecurityRole::Server)
.expect("server codec"),
);
let encrypted = client
.encrypt_frame(Frame::new(
1,
MsgType::RequestBody,
0,
Bytes::from_static(b"secret"),
))
.expect("encrypt");
let wire = Frame::decode(encrypted).expect("wire frame");
let barrier = Arc::new(Barrier::new(3));
let workers = (0..2)
.map(|_| {
let server = Arc::clone(&server);
let barrier = Arc::clone(&barrier);
let wire = wire.clone();
std::thread::spawn(move || {
barrier.wait();
server.decrypt_frame(wire)
})
})
.collect::<Vec<_>>();
barrier.wait();
let results = workers
.into_iter()
.map(|worker| worker.join().expect("worker should not panic"))
.collect::<Vec<_>>();
assert_eq!(results.iter().filter(|result| result.is_ok()).count(), 1);
assert_eq!(
results
.iter()
.filter(|result| matches!(result, Err(TunnelSecurityError::UnexpectedSequence)))
.count(),
1
);
}
#[test]
fn secure_frame_fails_closed_before_sequence_wrap() {
let codec = SecureFrameCodec::new(&test_key(), "session-1", TunnelSecurityRole::Client)
.expect("codec");
codec.next_sequence.store(u64::MAX, Ordering::Relaxed);
assert!(matches!(
codec.encrypt_frame(Frame::new(
1,
MsgType::RequestBody,
0,
Bytes::from_static(b"secret"),
)),
Err(TunnelSecurityError::SequenceExhausted)
));
}
#[test]
fn secure_frame_rejects_out_of_order_sequence_without_advancing() {
let client = SecureFrameCodec::new(&test_key(), "session-1", TunnelSecurityRole::Client)
@@ -335,4 +687,195 @@ mod tests {
assert_ne!(encrypted_a, encrypted_b);
}
#[test]
fn handshake_proof_round_trips_and_binds_every_field() {
let key = test_key();
let signature = sign_tunnel_security_handshake(
&key,
"node-1",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"0123456789abcdef0123456789abcdef",
3,
1_700_000_000,
"abcdef0123456789abcdef0123456789",
)
.expect("sign handshake");
assert_eq!(signature.len(), 43);
assert!(verify_tunnel_security_handshake(
&key,
"node-1",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"0123456789abcdef0123456789abcdef",
3,
1_700_000_000,
"abcdef0123456789abcdef0123456789",
&signature,
));
for valid in [
verify_tunnel_security_handshake(
&key,
"node-2",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"0123456789abcdef0123456789abcdef",
3,
1_700_000_000,
"abcdef0123456789abcdef0123456789",
&signature,
),
verify_tunnel_security_handshake(
&key,
"node-1",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"1123456789abcdef0123456789abcdef",
3,
1_700_000_000,
"abcdef0123456789abcdef0123456789",
&signature,
),
verify_tunnel_security_handshake(
&key,
"node-1",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"0123456789abcdef0123456789abcdef",
2,
1_700_000_000,
"abcdef0123456789abcdef0123456789",
&signature,
),
verify_tunnel_security_handshake(
&key,
"node-1",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"0123456789abcdef0123456789abcdef",
3,
1_700_000_001,
"abcdef0123456789abcdef0123456789",
&signature,
),
verify_tunnel_security_handshake(
&key,
"node-1",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"0123456789abcdef0123456789abcdef",
3,
1_700_000_000,
"bbcdef0123456789abcdef0123456789",
&signature,
),
] {
assert!(!valid, "tampered handshake field must fail verification");
}
}
#[test]
fn handshake_proof_rejects_wrong_key_and_malformed_signature() {
let key = test_key();
let signature = sign_tunnel_security_handshake(
&key,
"node-1",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"0123456789abcdef0123456789abcdef",
3,
1_700_000_000,
"abcdef0123456789abcdef0123456789",
)
.expect("sign handshake");
let wrong_key = base64::engine::general_purpose::STANDARD.encode([8_u8; 32]);
assert!(!verify_tunnel_security_handshake(
&wrong_key,
"node-1",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"0123456789abcdef0123456789abcdef",
3,
1_700_000_000,
"abcdef0123456789abcdef0123456789",
&signature,
));
assert!(!verify_tunnel_security_handshake(
&key,
"node-1",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"0123456789abcdef0123456789abcdef",
3,
1_700_000_000,
"abcdef0123456789abcdef0123456789",
"not-base64",
));
}
#[test]
fn control_plane_proof_round_trips_and_binds_identity_route_and_body() {
let key = test_key();
let body = br#"{"node_id":"node-1","heartbeat_id":7}"#;
let signature = sign_tunnel_control_plane_request(
&key,
"POST",
"/api/internal/tunnel/heartbeat",
"node-1",
1_700_000_000,
"nonce-0123456789abcdef",
body,
)
.expect("sign control-plane request");
assert!(verify_tunnel_control_plane_request(
&key,
"POST",
"/api/internal/tunnel/heartbeat",
"node-1",
1_700_000_000,
"nonce-0123456789abcdef",
body,
&signature,
));
for valid in [
verify_tunnel_control_plane_request(
&key,
"GET",
"/api/internal/tunnel/heartbeat",
"node-1",
1_700_000_000,
"nonce-0123456789abcdef",
body,
&signature,
),
verify_tunnel_control_plane_request(
&key,
"POST",
"/api/internal/tunnel/node-status",
"node-1",
1_700_000_000,
"nonce-0123456789abcdef",
body,
&signature,
),
verify_tunnel_control_plane_request(
&key,
"POST",
"/api/internal/tunnel/heartbeat",
"node-2",
1_700_000_000,
"nonce-0123456789abcdef",
body,
&signature,
),
verify_tunnel_control_plane_request(
&key,
"POST",
"/api/internal/tunnel/heartbeat",
"node-1",
1_700_000_000,
"nonce-0123456789abcdef",
br#"{"node_id":"node-1","heartbeat_id":8}"#,
&signature,
),
] {
assert!(
!valid,
"tampered control-plane field must fail verification"
);
}
}
}