Bounded Debugging + Websocket Carriers + Carriers Negotiation

Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
This commit is contained in:
Alexey
2026-08-26 17:00:20 +03:00
parent 43cd84aaa5
commit 923c79796a
52 changed files with 3450 additions and 980 deletions
+99
View File
@@ -0,0 +1,99 @@
use super::*;
#[tokio::test]
async fn windows_restricted_webview_empty_cookie_preserves_the_carrier_flow() {
for (index, carrier) in [WebCarrier::Https, WebCarrier::HttpsLanes]
.into_iter()
.enumerate()
{
let capability = [12 + index as u8; 32];
let mut config = runtime_config(capability, carrier);
config.web.timeouts.long_poll_secs = 0;
let generation = test_runtime_generation(1, config);
let active_runtime = Arc::new(ArcSwap::from(Arc::clone(&generation)));
let runtime = WebProcessRuntime::start(active_runtime);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(capability);
let root = format!(
"GET /?bridge={encoded} HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nCookie:\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let root_response = request(&listener, &runtime, root).await;
let (root_headers, root_body) = split_response(&root_response);
assert!(root_headers.starts_with(b"HTTP/1.1 200"));
let root_body = std::str::from_utf8(root_body).unwrap();
let bootstrap = root_body
.split_once("bootstrap=\"")
.and_then(|(_, suffix)| suffix.split_once('"'))
.map(|(token, _)| token)
.unwrap();
let hello = frame::encode(FrameType::Hello, 0, &[1]);
let create = |cookie: &str| {
let mut request = format!(
"POST /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {bootstrap}\r\nContent-Type: application/octet-stream\r\n{cookie}Content-Length: {}\r\nConnection: close\r\n\r\n",
hello.len()
)
.into_bytes();
request.extend_from_slice(&hello);
request
};
let nonempty_cookie =
request(&listener, &runtime, create("Cookie: state=unexpected\r\n")).await;
assert!(!nonempty_cookie.starts_with(b"HTTP/1.1 200"));
let duplicate_cookie = request(
&listener,
&runtime,
create("Cookie:\r\nCookie: state=unexpected\r\n"),
)
.await;
assert!(!duplicate_cookie.starts_with(b"HTTP/1.1 200"));
let create_response = request(&listener, &runtime, create("Cookie:\r\n")).await;
let (create_headers, create_body) = split_response(&create_response);
assert!(create_headers.starts_with(b"HTTP/1.1 200"));
assert_eq!(
response_header(create_headers, "x-carrier-mode"),
carrier.as_str()
);
assert_eq!(create_body, frame::encode(FrameType::Welcome, 0, &[]));
let session = response_header(create_headers, "x-session-token").to_string();
let lane = if carrier == WebCarrier::HttpsLanes {
"X-Lane-ID: 0\r\n"
} else {
""
};
let pong = frame::encode(FrameType::Pong, 0, &[]);
let mut uplink = format!(
"POST /api/v1/up HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nContent-Type: application/octet-stream\r\nCookie:\r\nX-Up-Seq: 1\r\n{lane}Content-Length: {}\r\nConnection: close\r\n\r\n",
pong.len()
)
.into_bytes();
uplink.extend_from_slice(&pong);
let uplink_response = request(&listener, &runtime, uplink).await;
let (uplink_headers, _) = split_response(&uplink_response);
assert!(uplink_headers.starts_with(b"HTTP/1.1 204"));
assert_eq!(response_header(uplink_headers, "x-up-ack"), "1");
let downlink = format!(
"POST /api/v1/down HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nCookie:\r\nX-Down-Cursor: 0\r\n{lane}Content-Length: 0\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let downlink_response = request(&listener, &runtime, downlink).await;
let (downlink_headers, _) = split_response(&downlink_response);
assert!(downlink_headers.starts_with(b"HTTP/1.1 204"));
assert_eq!(response_header(downlink_headers, "x-down-cursor"), "0");
let close = format!(
"DELETE /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nCookie:\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let close_response = request(&listener, &runtime, close).await;
assert!(close_response.starts_with(b"HTTP/1.1 204"));
runtime.shutdown().await;
generation.stop_sessions().await;
generation.stop_background_tasks().await;
}
}
+176
View File
@@ -0,0 +1,176 @@
use super::*;
use sha2::{Digest, Sha256};
const CAPABILITIES: &str = "https,https-lanes,websocket,websocket-lanes";
fn issue_bootstrap(runtime: &Arc<WebProcessRuntime>, client_ip: &str) -> String {
let profile = runtime
.active_generation()
.config()
.web
.runtime
.as_ref()
.unwrap()
.profiles[0]
.clone();
runtime
.issue_bootstrap(profile, client_ip.parse().unwrap())
.unwrap()
.token
}
fn create_request(
bootstrap: &str,
hello: &[u8],
attempt: Option<u8>,
failure: Option<&str>,
) -> Vec<u8> {
let negotiation = attempt.map_or_else(String::new, |attempt| {
let failure = failure
.map(|failure| format!("X-Carrier-Failure: {failure}\r\n"))
.unwrap_or_default();
format!(
"X-Carrier-Capabilities: {CAPABILITIES}\r\nX-Carrier-Attempt: {attempt}\r\n{failure}"
)
});
let mut request = format!(
"POST /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {bootstrap}\r\nContent-Type: application/octet-stream\r\n{negotiation}Content-Length: {}\r\nConnection: close\r\n\r\n",
hello.len()
)
.into_bytes();
request.extend_from_slice(hello);
request
}
fn optional_response_header<'a>(headers: &'a [u8], name: &str) -> Option<&'a str> {
std::str::from_utf8(headers)
.unwrap()
.lines()
.filter_map(|line| line.split_once(':'))
.find_map(|(header, value)| header.eq_ignore_ascii_case(name).then_some(value.trim()))
}
fn token_hash(token: &str) -> crate::web::manager::TokenHash {
let raw = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(token)
.unwrap();
Sha256::digest(raw).into()
}
#[tokio::test]
async fn absent_carriers_reject_negotiation_and_preserve_legacy_creation() {
let capability = [41; 32];
let generation = test_runtime_generation(1, runtime_config(capability, WebCarrier::Https));
let runtime = WebProcessRuntime::start(Arc::new(ArcSwap::from(Arc::clone(&generation))));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let bootstrap = issue_bootstrap(&runtime, "192.0.2.10");
let hello = frame::encode(FrameType::Hello, 0, &[1]);
let rejected = request(
&listener,
&runtime,
create_request(&bootstrap, &hello, Some(1), None),
)
.await;
let (rejected_headers, _) = split_response(&rejected);
assert!(optional_response_header(rejected_headers, "x-session-token").is_none());
let legacy = request(
&listener,
&runtime,
create_request(&bootstrap, &hello, None, None),
)
.await;
let (legacy_headers, _) = split_response(&legacy);
assert!(legacy_headers.starts_with(b"HTTP/1.1 200"));
assert_eq!(response_header(legacy_headers, "x-carrier-mode"), "https");
assert!(optional_response_header(legacy_headers, "x-carrier-attempt").is_none());
runtime.shutdown().await;
generation.stop_sessions().await;
generation.stop_background_tasks().await;
}
#[tokio::test]
async fn negotiation_replays_replaces_and_freezes_after_carrier_commit() {
let capability = [42; 32];
let config = negotiation_runtime_config(
capability,
WebCarrier::Websocket,
false,
Arc::from([
WebCarrier::Https,
WebCarrier::HttpsLanes,
WebCarrier::Websocket,
]),
);
let generation = test_runtime_generation(1, config);
let runtime = WebProcessRuntime::start(Arc::new(ArcSwap::from(Arc::clone(&generation))));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let bootstrap = issue_bootstrap(&runtime, "192.0.2.10");
let hello = frame::encode(FrameType::Hello, 0, &[1]);
let first_request = create_request(&bootstrap, &hello, Some(1), None);
let first = request(&listener, &runtime, first_request.clone()).await;
let (first_headers, _) = split_response(&first);
assert_eq!(response_header(first_headers, "x-carrier-mode"), "https");
assert_eq!(response_header(first_headers, "x-carrier-attempt"), "1");
let first_token = response_header(first_headers, "x-session-token").to_string();
let replay = request(&listener, &runtime, first_request).await;
let (replay_headers, _) = split_response(&replay);
assert_eq!(response_header(replay_headers, "x-session-token"), first_token);
let second_request = create_request(&bootstrap, &hello, Some(2), Some("timeout"));
let second = request(&listener, &runtime, second_request.clone()).await;
let (second_headers, _) = split_response(&second);
assert_eq!(
response_header(second_headers, "x-carrier-mode"),
"https-lanes"
);
assert_eq!(response_header(second_headers, "x-carrier-attempt"), "2");
let second_token = response_header(second_headers, "x-session-token").to_string();
assert_ne!(first_token, second_token);
assert!(
runtime
.get_session(token_hash(&first_token), "proxy.example.com")
.is_err()
);
let second_replay = request(&listener, &runtime, second_request).await;
let (second_replay_headers, _) = split_response(&second_replay);
assert_eq!(
response_header(second_replay_headers, "x-session-token"),
second_token
);
let open = frame::encode(FrameType::Open, 7, &[]);
let mut uplink = format!(
"POST /api/v1/up HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {second_token}\r\nContent-Type: application/octet-stream\r\nX-Up-Seq: 1\r\nX-Lane-ID: 7\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
open.len()
)
.into_bytes();
uplink.extend_from_slice(&open);
let committed = request(&listener, &runtime, uplink).await;
assert!(committed.starts_with(b"HTTP/1.1 204"));
let third = request(
&listener,
&runtime,
create_request(&bootstrap, &hello, Some(3), Some("http")),
)
.await;
let (third_headers, _) = split_response(&third);
assert!(optional_response_header(third_headers, "x-session-token").is_none());
assert!(
runtime
.get_session(token_hash(&second_token), "proxy.example.com")
.unwrap()
.is_carrier_committed()
);
runtime.shutdown().await;
generation.stop_sessions().await;
generation.stop_background_tasks().await;
}
+196 -1
View File
@@ -9,7 +9,11 @@ use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;
use crate::config::{WebClientIpSource, WebRuntimeProfile, WebRuntimeVhost};
use crate::web::manager::TokenHash;
use crate::web::manager::{
CarrierCapabilities, CarrierClientClass, CarrierFailure, CarrierRequest, TokenHash,
};
const USER_AGENT_CONTEXT: &[u8] = b"telemt-web-carrier-user-agent-v1\0";
/// Parses one lowercase canonical Host value restricted to the public HTTPS port.
pub(super) fn canonical_request_host<B>(request: &Request<B>) -> Option<&str> {
@@ -167,6 +171,124 @@ pub(super) fn canonical_u64_header<B>(request: &Request<B>, name: &'static str)
(parsed.to_string() == value).then_some(parsed)
}
/// Parses strict bridge negotiation metadata without trusting it for authentication.
pub(super) fn carrier_request<B>(request: &Request<B>, host: &str) -> Option<CarrierRequest> {
let user_agent_hash = normalized_user_agent_hash(request)?;
let capabilities = single_header(request, "x-carrier-capabilities");
let attempt = optional_canonical_u8_header(request, "x-carrier-attempt")?;
let failure = optional_failure_header(request)?;
match (capabilities, attempt) {
(None, None) if failure.is_none() => Some(CarrierRequest::legacy(user_agent_hash)),
(Some(capabilities), Some(attempt)) => {
let capabilities = parse_capabilities(capabilities)?;
if (attempt == 1) != failure.is_none() {
return None;
}
Some(CarrierRequest::automatic(
CarrierClientClass::Bridge,
capabilities,
attempt,
failure,
user_agent_hash,
))
}
(None, Some(attempt)) if strict_browser_hint(request, host) => {
if (attempt == 1) != failure.is_none() {
return None;
}
Some(CarrierRequest::automatic(
CarrierClientClass::BrowserHint,
CarrierCapabilities::all(),
attempt,
failure,
user_agent_hash,
))
}
_ => None,
}
}
fn parse_capabilities(value: &str) -> Option<CarrierCapabilities> {
let mut bits = 0u8;
let mut previous = None;
for token in value.split(',') {
let index = match token {
"https" => 0,
"https-lanes" => 1,
"websocket" => 2,
"websocket-lanes" => 3,
_ => return None,
};
if previous.is_some_and(|previous| index <= previous) {
return None;
}
previous = Some(index);
bits |= 1 << index;
}
CarrierCapabilities::from_bits(bits)
}
fn strict_browser_hint<B>(request: &Request<B>, host: &str) -> bool {
single_header(request, header::ORIGIN)
.is_some_and(|value| value == format!("https://{host}"))
&& single_header(request, "sec-fetch-site") == Some("same-origin")
&& single_header(request, "sec-fetch-mode") == Some("cors")
&& single_header(request, "sec-fetch-dest") == Some("empty")
}
fn optional_canonical_u8_header<B>(
request: &Request<B>,
name: &'static str,
) -> Option<Option<u8>> {
if !request.headers().contains_key(name) {
return Some(None);
}
canonical_u64_header(request, name)
.and_then(|value| u8::try_from(value).ok())
.filter(|value| (1..=4).contains(value))
.map(Some)
}
fn optional_failure_header<B>(request: &Request<B>) -> Option<Option<CarrierFailure>> {
if !request.headers().contains_key("x-carrier-failure") {
return Some(None);
}
single_header(request, "x-carrier-failure")
.and_then(CarrierFailure::parse)
.map(Some)
}
fn normalized_user_agent_hash<B>(request: &Request<B>) -> Option<[u8; 32]> {
let user_agent = match single_header(request, header::USER_AGENT) {
Some(value) => value.as_bytes(),
None if !request.headers().contains_key(header::USER_AGENT) => &[],
None => return None,
};
let mut digest = Sha256::new();
digest.update(USER_AGENT_CONTEXT);
let mut emitted = false;
let mut pending_whitespace = false;
for &byte in user_agent {
if byte.is_ascii_whitespace() {
pending_whitespace = emitted;
} else {
if pending_whitespace {
digest.update([b' ']);
}
digest.update([byte.to_ascii_lowercase()]);
emitted = true;
pending_whitespace = false;
}
}
Some(digest.finalize().into())
}
fn single_header<B>(request: &Request<B>, name: impl header::AsHeaderName) -> Option<&str> {
let mut values = request.headers().get_all(name).iter();
let value = values.next()?.to_str().ok()?;
values.next().is_none().then_some(value)
}
#[cfg(test)]
mod tests {
use super::*;
@@ -319,4 +441,77 @@ mod tests {
.append(header::COOKIE, "state=unexpected".parse().unwrap());
assert!(!compatible_cookie_header(&duplicate_mixed));
}
#[test]
fn carrier_metadata_is_canonical_and_legacy_safe() {
let automatic = Request::builder()
.header(
"x-carrier-capabilities",
"https,https-lanes,websocket,websocket-lanes",
)
.header("x-carrier-attempt", "2")
.header("x-carrier-failure", "timeout")
.header(header::USER_AGENT, "Example Browser")
.body(())
.unwrap();
let parsed = carrier_request(&automatic, "proxy.example.com").unwrap();
assert!(parsed.is_automatic());
assert_eq!(parsed.attempt(), Some(2));
assert_eq!(parsed.failure(), Some(CarrierFailure::Timeout));
let missing_failure = Request::builder()
.header(
"x-carrier-capabilities",
"https,https-lanes,websocket,websocket-lanes",
)
.header("x-carrier-attempt", "2")
.body(())
.unwrap();
assert!(carrier_request(&missing_failure, "proxy.example.com").is_none());
let legacy = Request::builder()
.header(header::USER_AGENT, "Native")
.body(())
.unwrap();
assert!(!carrier_request(&legacy, "proxy.example.com")
.unwrap()
.is_automatic());
let reordered = Request::builder()
.header("x-carrier-capabilities", "websocket,https")
.header("x-carrier-attempt", "1")
.body(())
.unwrap();
assert!(carrier_request(&reordered, "proxy.example.com").is_none());
}
#[test]
fn strict_browser_metadata_recovers_a_stripped_capability_marker() {
let request = Request::builder()
.header("x-carrier-attempt", "1")
.header(header::ORIGIN, "https://proxy.example.com")
.header("sec-fetch-site", "same-origin")
.header("sec-fetch-mode", "cors")
.header("sec-fetch-dest", "empty")
.body(())
.unwrap();
let parsed = carrier_request(&request, "proxy.example.com").unwrap();
assert_eq!(parsed.class(), CarrierClientClass::BrowserHint);
}
#[test]
fn user_agent_learning_key_is_case_and_whitespace_normalized() {
let first = Request::builder()
.header(header::USER_AGENT, " Example\t Browser ")
.body(())
.unwrap();
let second = Request::builder()
.header(header::USER_AGENT, "example browser")
.body(())
.unwrap();
assert_eq!(
normalized_user_agent_hash(&first),
normalized_user_agent_hash(&second)
);
}
}
+40 -104
View File
@@ -10,20 +10,48 @@ use tokio_util::sync::CancellationToken;
use super::serve_connection;
use crate::config::{
ProxyConfig, WebCarrier, WebClientIpSource, WebRuntimeConfig, WebRuntimeDecoy,
ProxyConfig, WebCarrier, WebCarriers, WebClientIpSource, WebRuntimeConfig, WebRuntimeDecoy,
WebRuntimeProfile, WebRuntimeVhost, WebSecretMode, WebStaticAsset, WebStaticSite,
};
use crate::maestro::generation::test_runtime_generation;
use crate::web::frame::{self, FrameType};
use crate::web::manager::WebProcessRuntime;
#[path = "legacy_tests.rs"]
mod legacy_tests;
#[path = "negotiation_tests.rs"]
mod negotiation_tests;
pub(super) fn runtime_config(capability: [u8; 32], carrier: WebCarrier) -> ProxyConfig {
runtime_config_with_carriers(capability, carrier, false, true, Arc::from([carrier]))
}
pub(super) fn negotiation_runtime_config(
capability: [u8; 32],
carrier: WebCarrier,
carrier_learning: bool,
carriers: Arc<[WebCarrier]>,
) -> ProxyConfig {
runtime_config_with_carriers(capability, carrier, true, carrier_learning, carriers)
}
fn runtime_config_with_carriers(
capability: [u8; 32],
carrier: WebCarrier,
carrier_negotiation_enabled: bool,
carrier_learning: bool,
carriers: Arc<[WebCarrier]>,
) -> ProxyConfig {
let profile = Arc::new(WebRuntimeProfile {
host: "proxy.example.com".to_string(),
public_addr: "203.0.113.10:443".parse().unwrap(),
user: "alice".to_string(),
secret_mode: WebSecretMode::Plain,
carrier,
carrier_negotiation_enabled,
carrier_learning,
carriers: Arc::clone(&carriers),
carrier_negotiation_deadlines_secs: [3, 5, 8, 12],
capability,
key_fingerprint: "0000000000000000".to_string(),
max_sessions: 4,
@@ -63,6 +91,12 @@ pub(super) fn runtime_config(capability: [u8; 32], carrier: WebCarrier) -> Proxy
let mut config = ProxyConfig::default();
config.web.enabled = true;
config.web.carrier = carrier;
config.web.carriers = if carrier_negotiation_enabled {
WebCarriers::Enabled(carriers.to_vec())
} else {
WebCarriers::Disabled
};
config.web.carrier_learning = carrier_learning;
config.web.limits.max_bootstraps_per_ip = 1;
config.web.timeouts.shutdown_secs = 1;
config.web.runtime = Some(Arc::new(WebRuntimeConfig {
@@ -106,7 +140,7 @@ pub(super) fn split_response(response: &[u8]) -> (&[u8], &[u8]) {
(&response[..separator], &response[separator + 4..])
}
fn response_header<'a>(headers: &'a [u8], name: &str) -> &'a str {
pub(super) fn response_header<'a>(headers: &'a [u8], name: &str) -> &'a str {
std::str::from_utf8(headers)
.unwrap()
.lines()
@@ -187,11 +221,9 @@ async fn https_carrier_bootstraps_and_closes_one_session() {
.windows(11)
.any(|value| value == b"bootstrap=\"")
);
assert!(
next_root_body
.windows(21)
.any(|value| value == b"carrier='https-lanes'")
);
assert!(next_root_body
.windows(b"const negotiationEnabled=false".len())
.any(|value| value == b"const negotiationEnabled=false"));
let close = format!(
"DELETE /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
@@ -375,7 +407,7 @@ async fn https_lanes_is_advertised_and_requires_canonical_lane_headers() {
let root_response = request(&listener, &runtime, root).await;
let (_, root_body) = split_response(&root_response);
let root_body = std::str::from_utf8(root_body).unwrap();
assert!(root_body.contains("carrier='https-lanes'"));
assert!(root_body.contains("const negotiationEnabled=false"));
let bootstrap = root_body
.split_once("bootstrap=\"")
.and_then(|(_, suffix)| suffix.split_once('"'))
@@ -437,99 +469,3 @@ async fn https_lanes_is_advertised_and_requires_canonical_lane_headers() {
generation.stop_sessions().await;
generation.stop_background_tasks().await;
}
#[tokio::test]
async fn windows_restricted_webview_empty_cookie_preserves_the_carrier_flow() {
for (index, carrier) in [WebCarrier::Https, WebCarrier::HttpsLanes]
.into_iter()
.enumerate()
{
let capability = [12 + index as u8; 32];
let mut config = runtime_config(capability, carrier);
config.web.timeouts.long_poll_secs = 0;
let generation = test_runtime_generation(1, config);
let active_runtime = Arc::new(ArcSwap::from(Arc::clone(&generation)));
let runtime = WebProcessRuntime::start(active_runtime);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(capability);
let root = format!(
"GET /?bridge={encoded} HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nCookie:\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let root_response = request(&listener, &runtime, root).await;
let (root_headers, root_body) = split_response(&root_response);
assert!(root_headers.starts_with(b"HTTP/1.1 200"));
let root_body = std::str::from_utf8(root_body).unwrap();
let bootstrap = root_body
.split_once("bootstrap=\"")
.and_then(|(_, suffix)| suffix.split_once('"'))
.map(|(token, _)| token)
.unwrap();
let hello = frame::encode(FrameType::Hello, 0, &[1]);
let create = |cookie: &str| {
let mut request = format!(
"POST /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {bootstrap}\r\nContent-Type: application/octet-stream\r\n{cookie}Content-Length: {}\r\nConnection: close\r\n\r\n",
hello.len()
)
.into_bytes();
request.extend_from_slice(&hello);
request
};
let nonempty_cookie =
request(&listener, &runtime, create("Cookie: state=unexpected\r\n")).await;
assert!(!nonempty_cookie.starts_with(b"HTTP/1.1 200"));
let duplicate_cookie = request(
&listener,
&runtime,
create("Cookie:\r\nCookie: state=unexpected\r\n"),
)
.await;
assert!(!duplicate_cookie.starts_with(b"HTTP/1.1 200"));
let create_response = request(&listener, &runtime, create("Cookie:\r\n")).await;
let (create_headers, create_body) = split_response(&create_response);
assert!(create_headers.starts_with(b"HTTP/1.1 200"));
assert_eq!(
response_header(create_headers, "x-carrier-mode"),
carrier.as_str()
);
assert_eq!(create_body, frame::encode(FrameType::Welcome, 0, &[]));
let session = response_header(create_headers, "x-session-token").to_string();
let lane = (carrier == WebCarrier::HttpsLanes)
.then_some("X-Lane-ID: 0\r\n")
.unwrap_or_default();
let pong = frame::encode(FrameType::Pong, 0, &[]);
let mut uplink = format!(
"POST /api/v1/up HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nContent-Type: application/octet-stream\r\nCookie:\r\nX-Up-Seq: 1\r\n{lane}Content-Length: {}\r\nConnection: close\r\n\r\n",
pong.len()
)
.into_bytes();
uplink.extend_from_slice(&pong);
let uplink_response = request(&listener, &runtime, uplink).await;
let (uplink_headers, _) = split_response(&uplink_response);
assert!(uplink_headers.starts_with(b"HTTP/1.1 204"));
assert_eq!(response_header(uplink_headers, "x-up-ack"), "1");
let downlink = format!(
"POST /api/v1/down HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nCookie:\r\nX-Down-Cursor: 0\r\n{lane}Content-Length: 0\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let downlink_response = request(&listener, &runtime, downlink).await;
let (downlink_headers, _) = split_response(&downlink_response);
assert!(downlink_headers.starts_with(b"HTTP/1.1 204"));
assert_eq!(response_header(downlink_headers, "x-down-cursor"), "0");
let close = format!(
"DELETE /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nCookie:\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let close_response = request(&listener, &runtime, close).await;
assert!(close_response.starts_with(b"HTTP/1.1 204"));
runtime.shutdown().await;
generation.stop_sessions().await;
generation.stop_background_tasks().await;
}
}
+32 -18
View File
@@ -200,7 +200,6 @@ impl AsyncRead for ConnectionIo {
Poll::Ready(Ok(())) => {
let filled = limited.filled().len();
boundary.observe(limited.filled());
drop(limited);
buffer.advance(filled);
Poll::Ready(Ok(()))
}
@@ -244,6 +243,7 @@ struct ParsedUpgrade {
protocol: String,
accept: String,
carrier: ParsedCarrier,
acknowledge_commit: bool,
}
pub(super) async fn handle(
@@ -325,6 +325,7 @@ pub(super) async fn handle(
connection,
lane_reservation.take(),
trace_context,
parsed.acknowledge_commit,
)
.await;
});
@@ -376,21 +377,18 @@ fn parse_upgrade<B>(request: &Request<B>) -> Option<ParsedUpgrade> {
{
return None;
}
let (token, carrier) = if let Some(token) = protocol.strip_prefix("tproxy-v1.") {
(token, ParsedCarrier::Multiplex)
let (token, carrier, acknowledge_commit) = if let Some(token) =
protocol.strip_prefix("tproxy-auto-v1.")
{
(token, ParsedCarrier::Multiplex, true)
} else if let Some(lane) = protocol.strip_prefix("tproxy-auto-lane-v1.") {
let (token, lane_id) = parse_lane_protocol(lane)?;
(token, ParsedCarrier::Lane(lane_id), true)
} else if let Some(token) = protocol.strip_prefix("tproxy-v1.") {
(token, ParsedCarrier::Multiplex, false)
} else if let Some(lane) = protocol.strip_prefix("tproxy-lane-v1.") {
let (token, lane_id) = lane.split_once('.')?;
if lane_id.is_empty()
|| lane_id.starts_with('+')
|| (lane_id.len() > 1 && lane_id.starts_with('0'))
{
return None;
}
let lane_id = lane_id
.parse::<u32>()
.ok()
.filter(|value| (1..=crate::web::frame::MAX_STREAM_ID).contains(value))?;
(token, ParsedCarrier::Lane(lane_id))
let (token, lane_id) = parse_lane_protocol(lane)?;
(token, ParsedCarrier::Lane(lane_id), false)
} else {
return None;
};
@@ -412,13 +410,29 @@ fn parse_upgrade<B>(request: &Request<B>) -> Option<ParsedUpgrade> {
protocol: protocol.to_string(),
accept: base64::engine::general_purpose::STANDARD.encode(accept.finalize()),
carrier,
acknowledge_commit,
})
}
fn single_header<'a, B>(
request: &'a Request<B>,
fn parse_lane_protocol(value: &str) -> Option<(&str, u32)> {
let (token, lane_id) = value.split_once('.')?;
if lane_id.is_empty()
|| lane_id.starts_with('+')
|| (lane_id.len() > 1 && lane_id.starts_with('0'))
{
return None;
}
let lane_id = lane_id
.parse::<u32>()
.ok()
.filter(|value| (1..=crate::web::frame::MAX_STREAM_ID).contains(value))?;
Some((token, lane_id))
}
fn single_header<B>(
request: &Request<B>,
name: impl hyper::header::AsHeaderName,
) -> Option<&'a str> {
) -> Option<&str> {
let mut values = request.headers().get_all(name).iter();
let value = values.next()?.to_str().ok()?;
values.next().is_none().then_some(value)
+36 -203
View File
@@ -2,7 +2,6 @@ use std::sync::Arc;
use std::time::{Duration, Instant};
use bytes::Bytes;
use futures_util::{SinkExt, StreamExt};
use hyper_util::rt::TokioIo;
use tokio_tungstenite::WebSocketStream;
use tokio_tungstenite::tungstenite::protocol::{Message, Role, WebSocketConfig};
@@ -10,7 +9,7 @@ use tokio_util::sync::CancellationToken;
use super::ConnectionIo;
use crate::web::manager::{
ManagerError, WebProcessRuntime, WebSocketBudgetLease, WebSocketConnection,
WebProcessRuntime, WebSocketBudgetLease, WebSocketConnection,
};
use crate::web::session::{WebSession, WebSocketLaneReservation};
use crate::web::trace::{TraceDirection, TraceWebSocketContext};
@@ -18,6 +17,10 @@ use crate::web::trace::{TraceDirection, TraceWebSocketContext};
const READ_BUFFER_BYTES: usize = 64 * 1024;
const WRITE_BUFFER_BYTES: usize = 64 * 1024;
// Cancellation-safe message I/O and budget retries remain separate from carrier loops.
mod io;
use io::{flush, process_lane, process_multiplex, read_message, record_message, reserve_data, send};
pub(super) async fn run_upgraded(
on_upgrade: hyper::upgrade::OnUpgrade,
runtime: Arc<WebProcessRuntime>,
@@ -25,6 +28,7 @@ pub(super) async fn run_upgraded(
connection: WebSocketConnection,
mut lane_reservation: Option<WebSocketLaneReservation>,
trace: Option<TraceWebSocketContext>,
acknowledge_commit: bool,
) {
let Ok(upgraded) = on_upgrade.await else {
return;
@@ -57,6 +61,7 @@ pub(super) async fn run_upgraded(
reservation,
cancellation.clone(),
trace.as_ref(),
acknowledge_commit,
)
.await;
} else {
@@ -67,6 +72,7 @@ pub(super) async fn run_upgraded(
&connection,
cancellation.clone(),
trace.as_ref(),
acknowledge_commit,
)
.await;
}
@@ -82,7 +88,7 @@ pub(super) async fn run_upgraded(
if let Some(reservation) = lane_reservation {
session.close_websocket_lane(reservation.lane_id());
drop(reservation);
} else {
} else if !acknowledge_commit || session.is_carrier_committed() {
session.close();
}
}
@@ -96,6 +102,7 @@ async fn run_multiplex(
connection: &WebSocketConnection,
cancellation: CancellationToken,
trace: Option<&TraceWebSocketContext>,
acknowledge_commit: bool,
) -> Result<(), ()> {
let mut sequence = 1u64;
let mut cursor = 0u64;
@@ -136,6 +143,18 @@ async fn run_multiplex(
started,
);
result?;
if acknowledge_commit && sequence == 1 && session.is_carrier_committed() {
let started = Instant::now();
send(socket, runtime, Message::Binary(Bytes::new())).await?;
record_message(
runtime,
trace,
TraceDirection::Response,
"carrier-ack",
&[],
started,
);
}
sequence = sequence.checked_add(1).ok_or(())?;
connection.mark_peer_activity();
next_ping = Instant::now() + liveness_interval;
@@ -255,6 +274,7 @@ async fn run_multiplex(
}
}
#[allow(clippy::too_many_arguments)]
async fn run_lane(
socket: &mut CarrierSocket,
runtime: &Arc<WebProcessRuntime>,
@@ -263,6 +283,7 @@ async fn run_lane(
reservation: &mut WebSocketLaneReservation,
cancellation: CancellationToken,
trace: Option<&TraceWebSocketContext>,
acknowledge_commit: bool,
) -> Result<(), ()> {
let mut sequence = 1u64;
let mut cursor = 0u64;
@@ -309,6 +330,18 @@ async fn run_lane(
started,
);
result?;
if acknowledge_commit && sequence == 1 && session.is_carrier_committed() {
let started = Instant::now();
send(socket, runtime, Message::Binary(Bytes::new())).await?;
record_message(
runtime,
trace,
TraceDirection::Response,
"carrier-ack",
&[],
started,
);
}
sequence = sequence.checked_add(1).ok_or(())?;
connection.mark_peer_activity();
next_ping = Instant::now() + liveness_interval;
@@ -436,203 +469,3 @@ enum DriverEvent {
Down(crate::web::session::PollResult),
Liveness,
}
async fn read_message(
socket: &mut CarrierSocket,
runtime: &Arc<WebProcessRuntime>,
owner: crate::web::manager::ProfileKey,
cancellation: &CancellationToken,
retained_budget: &mut Option<WebSocketBudgetLease>,
) -> Result<(Message, Option<WebSocketBudgetLease>), ()> {
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
ready = socket.get_ref().readable() => ready.map_err(|_| ())?,
}
if retained_budget.is_none() {
let maximum = runtime
.active_generation()
.config()
.web
.limits
.carrier_batch_bytes;
*retained_budget = Some(reserve_data(runtime, owner, maximum, cancellation).await?);
}
let message = tokio::select! {
_ = cancellation.cancelled() => return Err(()),
message = socket.next() => message.ok_or(())?.map_err(|_| ())?,
};
if socket.get_ref().websocket_fragmented_message() {
return Ok((message, None));
}
let mut budget = retained_budget.take().ok_or(())?;
budget.shrink_to(message.len());
Ok((message, Some(budget)))
}
async fn reserve_data(
runtime: &Arc<WebProcessRuntime>,
owner: crate::web::manager::ProfileKey,
bytes: usize,
cancellation: &CancellationToken,
) -> Result<WebSocketBudgetLease, ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_backpressure_secs,
);
tokio::time::timeout(timeout, async {
loop {
let notify = runtime.budget_notify();
let notified = notify.notified();
if let Some(budget) = runtime.try_websocket_data_budget(owner, bytes.max(1)) {
return Ok(budget);
}
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = notified => {}
}
}
})
.await
.map_err(|_| ())?
}
async fn process_multiplex(
runtime: &Arc<WebProcessRuntime>,
session: &Arc<WebSession>,
sequence: u64,
body: &[u8],
cancellation: &CancellationToken,
) -> Result<(), ()> {
retry_backpressure(runtime, cancellation, || {
session.process_up(sequence, body).map(|_| ())
})
.await
}
async fn process_lane(
runtime: &Arc<WebProcessRuntime>,
session: &Arc<WebSession>,
reservation: &mut WebSocketLaneReservation,
sequence: u64,
body: &[u8],
cancellation: &CancellationToken,
) -> Result<(), ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_backpressure_secs,
);
tokio::time::timeout(timeout, async {
loop {
let notify = runtime.budget_notify();
let notified = notify.notified();
match session.process_websocket_lane(reservation, sequence, body) {
Ok(()) => return Ok(()),
Err(ManagerError::Backpressure) => {}
Err(_) => return Err(()),
}
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = notified => {}
}
}
})
.await
.map_err(|_| ())?
}
async fn retry_backpressure<F>(
runtime: &Arc<WebProcessRuntime>,
cancellation: &CancellationToken,
mut operation: F,
) -> Result<(), ()>
where
F: FnMut() -> Result<(), ManagerError>,
{
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_backpressure_secs,
);
tokio::time::timeout(timeout, async {
loop {
let notify = runtime.budget_notify();
let notified = notify.notified();
match operation() {
Ok(()) => return Ok(()),
Err(ManagerError::Backpressure) => {}
Err(_) => return Err(()),
}
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = notified => {}
}
}
})
.await
.map_err(|_| ())?
}
async fn send(
socket: &mut CarrierSocket,
runtime: &WebProcessRuntime,
message: Message,
) -> Result<(), ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_write_secs,
);
tokio::time::timeout(timeout, socket.send(message))
.await
.map_err(|_| ())?
.map_err(|_| ())
}
async fn flush(socket: &mut CarrierSocket, runtime: &WebProcessRuntime) -> Result<(), ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_write_secs,
);
tokio::time::timeout(timeout, socket.flush())
.await
.map_err(|_| ())?
.map_err(|_| ())
}
fn record_message(
runtime: &WebProcessRuntime,
trace: Option<&TraceWebSocketContext>,
direction: TraceDirection,
message_type: &'static str,
payload: &[u8],
started: Instant,
) {
let Some(trace) = trace else {
return;
};
runtime.trace().record_websocket_message(
trace,
direction,
message_type,
payload,
started.elapsed().as_micros().min(u128::from(u64::MAX)) as u64,
);
}
+214
View File
@@ -0,0 +1,214 @@
use std::sync::Arc;
use std::time::{Duration, Instant};
use futures_util::{SinkExt, StreamExt};
use tokio_tungstenite::tungstenite::protocol::Message;
use tokio_util::sync::CancellationToken;
use super::CarrierSocket;
use crate::web::manager::{ManagerError, WebProcessRuntime, WebSocketBudgetLease};
use crate::web::session::{WebSession, WebSocketLaneReservation};
use crate::web::trace::{TraceDirection, TraceWebSocketContext};
pub(super) async fn read_message(
socket: &mut CarrierSocket,
runtime: &Arc<WebProcessRuntime>,
owner: crate::web::manager::ProfileKey,
cancellation: &CancellationToken,
retained_budget: &mut Option<WebSocketBudgetLease>,
) -> Result<(Message, Option<WebSocketBudgetLease>), ()> {
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
ready = socket.get_ref().readable() => ready.map_err(|_| ())?,
}
if retained_budget.is_none() {
let maximum = runtime
.active_generation()
.config()
.web
.limits
.carrier_batch_bytes;
*retained_budget = Some(reserve_data(runtime, owner, maximum, cancellation).await?);
}
let message = tokio::select! {
_ = cancellation.cancelled() => return Err(()),
message = socket.next() => message.ok_or(())?.map_err(|_| ())?,
};
if socket.get_ref().websocket_fragmented_message() {
return Ok((message, None));
}
let mut budget = retained_budget.take().ok_or(())?;
budget.shrink_to(message.len());
Ok((message, Some(budget)))
}
pub(super) async fn reserve_data(
runtime: &Arc<WebProcessRuntime>,
owner: crate::web::manager::ProfileKey,
bytes: usize,
cancellation: &CancellationToken,
) -> Result<WebSocketBudgetLease, ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_backpressure_secs,
);
tokio::time::timeout(timeout, async {
loop {
let notify = runtime.budget_notify();
let notified = notify.notified();
if let Some(budget) = runtime.try_websocket_data_budget(owner, bytes.max(1)) {
return Ok(budget);
}
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = notified => {}
}
}
})
.await
.map_err(|_| ())?
}
pub(super) async fn process_multiplex(
runtime: &Arc<WebProcessRuntime>,
session: &Arc<WebSession>,
sequence: u64,
body: &[u8],
cancellation: &CancellationToken,
) -> Result<(), ()> {
retry_backpressure(runtime, cancellation, || {
session.process_up(sequence, body).map(|_| ())
})
.await
}
pub(super) async fn process_lane(
runtime: &Arc<WebProcessRuntime>,
session: &Arc<WebSession>,
reservation: &mut WebSocketLaneReservation,
sequence: u64,
body: &[u8],
cancellation: &CancellationToken,
) -> Result<(), ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_backpressure_secs,
);
tokio::time::timeout(timeout, async {
loop {
let notify = runtime.budget_notify();
let notified = notify.notified();
match session.process_websocket_lane(reservation, sequence, body) {
Ok(()) => return Ok(()),
Err(ManagerError::Backpressure) => {}
Err(_) => return Err(()),
}
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = notified => {}
}
}
})
.await
.map_err(|_| ())?
}
async fn retry_backpressure<F>(
runtime: &Arc<WebProcessRuntime>,
cancellation: &CancellationToken,
mut operation: F,
) -> Result<(), ()>
where
F: FnMut() -> Result<(), ManagerError>,
{
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_backpressure_secs,
);
tokio::time::timeout(timeout, async {
loop {
let notify = runtime.budget_notify();
let notified = notify.notified();
match operation() {
Ok(()) => return Ok(()),
Err(ManagerError::Backpressure) => {}
Err(_) => return Err(()),
}
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = notified => {}
}
}
})
.await
.map_err(|_| ())?
}
pub(super) async fn send(
socket: &mut CarrierSocket,
runtime: &WebProcessRuntime,
message: Message,
) -> Result<(), ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_write_secs,
);
tokio::time::timeout(timeout, socket.send(message))
.await
.map_err(|_| ())?
.map_err(|_| ())
}
pub(super) async fn flush(
socket: &mut CarrierSocket,
runtime: &WebProcessRuntime,
) -> Result<(), ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_write_secs,
);
tokio::time::timeout(timeout, socket.flush())
.await
.map_err(|_| ())?
.map_err(|_| ())
}
pub(super) fn record_message(
runtime: &WebProcessRuntime,
trace: Option<&TraceWebSocketContext>,
direction: TraceDirection,
message_type: &'static str,
payload: &[u8],
started: Instant,
) {
let Some(trace) = trace else {
return;
};
runtime.trace().record_websocket_message(
trace,
direction,
message_type,
payload,
started.elapsed().as_micros().min(u128::from(u64::MAX)) as u64,
);
}
+150 -4
View File
@@ -14,8 +14,10 @@ use tokio_util::sync::CancellationToken;
use crate::maestro::generation::{RuntimeGeneration, test_runtime_generation};
use crate::web::frame::{self, FrameType};
use crate::web::http::tests::runtime_config;
use crate::web::manager::WebProcessRuntime;
use crate::web::http::tests::{negotiation_runtime_config, runtime_config};
use crate::web::manager::{
CarrierCapabilities, CarrierClientClass, CarrierFailure, CarrierRequest, WebProcessRuntime,
};
fn request(protocol: &str) -> Request<()> {
Request::builder()
@@ -39,6 +41,14 @@ fn canonical_multiplex_and_lane_protocols_are_accepted() {
let lane = parse_upgrade(&request(&format!("tproxy-lane-v1.{token}.16777215"))).unwrap();
assert!(matches!(lane.carrier, ParsedCarrier::Lane(16_777_215)));
let automatic = parse_upgrade(&request(&format!("tproxy-auto-v1.{token}"))).unwrap();
assert!(automatic.acknowledge_commit);
let automatic_lane =
parse_upgrade(&request(&format!("tproxy-auto-lane-v1.{token}.7"))).unwrap();
assert!(matches!(automatic_lane.carrier, ParsedCarrier::Lane(7)));
assert!(automatic_lane.acknowledge_commit);
assert!(!multiplex.acknowledge_commit);
}
#[test]
@@ -92,7 +102,20 @@ fn live_runtime(carrier: WebCarrier) -> LiveRuntime {
}
fn live_runtime_with_long_poll(carrier: WebCarrier, long_poll_secs: u64) -> LiveRuntime {
let mut config = runtime_config([31; 32], carrier);
live_runtime_from_config(runtime_config([31; 32], carrier), long_poll_secs)
}
fn live_negotiation_runtime(carrier: WebCarrier, carriers: Arc<[WebCarrier]>) -> LiveRuntime {
live_runtime_from_config(
negotiation_runtime_config([31; 32], carrier, false, carriers),
1,
)
}
fn live_runtime_from_config(
mut config: crate::config::ProxyConfig,
long_poll_secs: u64,
) -> LiveRuntime {
config.web.timeouts.long_poll_secs = long_poll_secs;
config.web.timeouts.websocket_write_secs = 2;
config.web.timeouts.websocket_backpressure_secs = 2;
@@ -105,6 +128,48 @@ fn live_runtime_with_long_poll(carrier: WebCarrier, long_poll_secs: u64) -> Live
}
}
fn create_automatic_session(
runtime: &Arc<WebProcessRuntime>,
) -> (TokenHash, Bytes, String, TokenHash) {
let profile = runtime
.active_generation()
.config()
.web
.runtime
.as_ref()
.unwrap()
.profiles[0]
.clone();
let client_ip = "192.0.2.10".parse().unwrap();
let bootstrap = runtime.issue_bootstrap(profile, client_ip).unwrap().token;
let raw = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(&bootstrap)
.unwrap();
let bootstrap_hash = Sha256::digest(raw).into();
let hello = frame::encode(FrameType::Hello, 0, &[1]);
let session = runtime
.create_session(
bootstrap_hash,
"proxy.example.com",
client_ip,
&hello,
CarrierRequest::automatic(
CarrierClientClass::Bridge,
CarrierCapabilities::all(),
1,
None,
[9; 32],
),
)
.unwrap()
.token;
let raw = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(&session)
.unwrap();
let session_hash = Sha256::digest(raw).into();
(bootstrap_hash, hello, session, session_hash)
}
fn create_session(runtime: &Arc<WebProcessRuntime>) -> (String, TokenHash) {
let profile = runtime
.active_generation()
@@ -123,7 +188,13 @@ fn create_session(runtime: &Arc<WebProcessRuntime>) -> (String, TokenHash) {
let bootstrap_hash = Sha256::digest(raw).into();
let hello = frame::encode(FrameType::Hello, 0, &[1]);
let session = runtime
.create_session(bootstrap_hash, "proxy.example.com", client_ip, &hello)
.create_session(
bootstrap_hash,
"proxy.example.com",
client_ip,
&hello,
CarrierRequest::legacy([0; 32]),
)
.unwrap()
.token;
let raw = base64::engine::general_purpose::URL_SAFE_NO_PAD
@@ -309,3 +380,78 @@ async fn malformed_websocket_lane_closes_only_that_lane() {
let _ = second.close(None).await;
live.shutdown().await;
}
#[tokio::test]
async fn automatic_websocket_carriers_ack_the_first_committing_message() {
for carrier in [WebCarrier::Websocket, WebCarrier::WebsocketLanes] {
let live = live_negotiation_runtime(carrier, Arc::from([carrier]));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let (_, _, session, session_hash) = create_automatic_session(&live.runtime);
let protocol = match carrier {
WebCarrier::Websocket => format!("tproxy-auto-v1.{session}"),
WebCarrier::WebsocketLanes => format!("tproxy-auto-lane-v1.{session}.7"),
_ => unreachable!(),
};
let mut socket = upgrade(&listener, &live.runtime, &protocol).await;
socket
.send(Message::Binary(frame::encode(FrameType::Open, 7, &[])))
.await
.unwrap();
let acknowledgement = tokio::time::timeout(Duration::from_secs(2), socket.next())
.await
.unwrap()
.unwrap()
.unwrap();
assert_eq!(acknowledgement, Message::Binary(Bytes::new()));
assert!(
live.runtime
.get_session(session_hash, "proxy.example.com")
.unwrap()
.is_carrier_committed()
);
let _ = socket.close(None).await;
live.shutdown().await;
}
}
#[tokio::test]
async fn failed_automatic_multiplex_socket_remains_supersedable() {
let live = live_negotiation_runtime(
WebCarrier::Https,
Arc::from([WebCarrier::Websocket, WebCarrier::Https]),
);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let (bootstrap_hash, hello, session, session_hash) =
create_automatic_session(&live.runtime);
let protocol = format!("tproxy-auto-v1.{session}");
let mut socket = upgrade(&listener, &live.runtime, &protocol).await;
socket.close(None).await.unwrap();
tokio::time::sleep(Duration::from_millis(200)).await;
assert!(
live.runtime
.get_session(session_hash, "proxy.example.com")
.is_ok()
);
let replacement = live
.runtime
.create_session(
bootstrap_hash,
"proxy.example.com",
"192.0.2.10".parse().unwrap(),
&hello,
CarrierRequest::automatic(
CarrierClientClass::Bridge,
CarrierCapabilities::all(),
2,
Some(CarrierFailure::Upgrade),
[9; 32],
),
)
.unwrap();
assert_eq!(replacement.carrier, WebCarrier::Https);
assert_eq!(replacement.attempt, Some(2));
live.shutdown().await;
}