use super::*;
#[tokio::test]
async fn queued_client_requests_do_not_starve_gateway_broadcasts() {
let root = tempfile::tempdir().unwrap();
let (server, grant) = configured_test_server(root.path().join("state")).await;
let identity = server.auth.pair(&grant.code, "busy client").unwrap();
let (revocations, _) = broadcast::channel(1);
let connection = ConnectionContext {
local: true,
auth: Arc::clone(&server.auth),
host: server.host.clone(),
bots: Arc::clone(&server.bots),
client_connections: Arc::new(ClientConnections::default()),
client_revocations: revocations,
admission: ConnectionAdmission::new(1, 1).admit().await,
access_lease: None,
};
let (client, stream) = tokio::io::duplex(1024 * 1024);
let (reader, mut writer) = tokio::io::split(client);
let mut reader = FrameReader::new(reader);
write_frame(
&mut writer,
&ClientFrame::new(ClientMessage::Authenticate {
token: identity.token,
client_kind: ClientKind::Cli,
catalog: Default::default(),
}),
)
.await
.unwrap();
let serving = serve_connection(stream, connection, Instant::now() + PRE_AUTH_TIMEOUT, None);
tokio::pin!(serving);
tokio::select! {
result = &mut serving => panic!("connection stopped: {result:?}"),
() = async {
while !matches!(read_frame::<ServerFrame>(&mut reader).await.unwrap().unwrap().message, ServerMessage::Ready { .. }) {}
} => {}
}
for id in 0..128 {
write_frame(
&mut writer,
&ClientFrame::new(ClientMessage::SetNotifications {
request_id: id.to_string(),
disabled: BTreeSet::new(),
}),
)
.await
.unwrap();
}
server
.host
.create_bot("Fairness", "Transport test")
.await
.unwrap();
tokio::select! {
result = &mut serving => panic!("connection stopped: {result:?}"),
() = async {
let mut responses = 0;
loop {
match read_frame::<ServerFrame>(&mut reader).await.unwrap().unwrap().message {
ServerMessage::Bots { .. } => break,
ServerMessage::Accepted { .. } => responses += 1,
_ => {}
}
assert!(responses < 128, "input drained before an already-ready broadcast");
}
} => {}
}
server.host.shutdown().await;
}
#[tokio::test(start_paused = true)]
async fn non_reading_client_releases_its_authenticated_connection_slot() {
let root = tempfile::tempdir().unwrap();
let (server, grant) = configured_test_server(root.path().join("state")).await;
let identity = server.auth.pair(&grant.code, "stalled").unwrap();
let admission = ConnectionAdmission::new(1, 1);
let (revocations, _) = broadcast::channel(1);
let connection = ConnectionContext {
local: true,
auth: Arc::clone(&server.auth),
host: server.host.clone(),
bots: Arc::clone(&server.bots),
client_connections: Arc::new(ClientConnections::default()),
client_revocations: revocations,
admission: admission.admit().await,
access_lease: None,
};
let (mut client, stream) = tokio::io::duplex(64);
let serving = tokio::spawn(serve_connection(
stream,
connection,
Instant::now() + PRE_AUTH_TIMEOUT,
None,
));
write_frame(
&mut client,
&ClientFrame::new(ClientMessage::Authenticate {
token: identity.token,
client_kind: ClientKind::Cli,
catalog: Default::default(),
}),
)
.await
.unwrap();
let error = serving
.await
.unwrap()
.expect_err("stalled write closes the connection");
assert!(matches!(error, Error::Io(error) if error.kind() == std::io::ErrorKind::TimedOut));
assert!(admission.admit().await.promote().is_some());
server.host.shutdown().await;
}
#[tokio::test]
async fn readiness_follows_initialization_and_allows_authenticated_connections() {
let root = tempfile::tempdir().unwrap();
let (mut server, grant) =
GatewayServer::bootstrap(root.path().join("state"), "127.0.0.1:0".parse().unwrap())
.await
.unwrap();
let endpoint: Endpoint = format!("tcp://{}", server.listen_addr()).parse().unwrap();
let mut ready = server.notify_ready();
assert!(matches!(
ready.try_recv(),
Err(tokio::sync::oneshot::error::TryRecvError::Empty)
));
let (stop, stopped) = tokio::sync::oneshot::channel();
let serving = tokio::spawn(server.serve_until(async {
let _ = stopped.await;
}));
ready.await.expect("serving loop ready");
let (client, _) = GatewayClient::pair(&endpoint, grant.code, "readiness", ClientKind::Cli)
.await
.unwrap();
let (_, mut events) = client.into_parts();
wait_gateway_ready(&mut events).await;
stop.send(()).unwrap();
serving.await.unwrap().unwrap();
}
#[tokio::test]
async fn failed_listener_initialization_never_signals_readiness() {
let root = tempfile::tempdir().unwrap();
let (mut server, _) =
GatewayServer::bootstrap(root.path().join("state"), "127.0.0.1:0".parse().unwrap())
.await
.unwrap();
server.config.tls = Some(TlsConfig {
certificate: root.path().join("missing.pem"),
private_key: root.path().join("missing.key"),
});
let ready = server.notify_ready();
assert!(server.serve_until(std::future::pending()).await.is_err());
assert!(ready.await.is_err(), "failure must not announce readiness");
}
#[test]
fn cloud_gateway_access_lease_fails_closed_after_five_minutes() {
let now = UNIX_EPOCH + Duration::from_secs(1_700_000_000);
assert!(access_lease(None, true, now).is_err());
assert!(access_lease(Some("1699999700"), true, now).is_err());
assert!(access_lease(Some("1699999701"), true, now).is_ok());
assert!(
access_lease(None, false, now)
.expect("local gateway")
.is_none()
);
}
#[tokio::test(start_paused = true)]
async fn access_lease_closes_existing_client_and_listener() {
let root = tempfile::tempdir().expect("temporary directory");
let (mut server, grant) = GatewayServer::bootstrap(
root.path().join("state"),
std::net::SocketAddr::from(([127, 0, 0, 1], 0)),
)
.await
.expect("bootstrap gateway");
let listen = server.listen_addr();
server.access_lease = Some(AccessLease {
expires_at: SystemTime::now() + Duration::from_secs(30),
deadline: Instant::now() + Duration::from_secs(30),
});
let serving = tokio::spawn(server.serve_until(std::future::pending()));
let endpoint = format!("tcp://{listen}")
.parse::<Endpoint>()
.expect("endpoint");
let (client, _) = GatewayClient::pair(&endpoint, grant.code, "lease test", ClientKind::Ios)
.await
.expect("pair client");
let (_, mut events) = client.into_parts();
wait_gateway_ready(&mut events).await;
tokio::time::advance(Duration::from_secs(30)).await;
serving
.await
.expect("gateway task")
.expect("lease shutdown");
assert!(events.next().await.expect("client disconnect").is_none());
assert!(TcpStream::connect(listen).await.is_err());
}
#[tokio::test(start_paused = true)]
async fn wall_clock_expiry_closes_a_warm_gateway_before_its_timer() {
let root = tempfile::tempdir().expect("temporary directory");
let (mut server, _) = GatewayServer::bootstrap(
root.path().join("state"),
std::net::SocketAddr::from(([127, 0, 0, 1], 0)),
)
.await
.expect("bootstrap gateway");
let listen = server.listen_addr();
server.access_lease = Some(AccessLease {
expires_at: SystemTime::now() - Duration::from_secs(1),
deadline: Instant::now() + Duration::from_secs(60),
});
server
.serve_until(std::future::pending())
.await
.expect("wall clock lease shutdown");
assert!(TcpStream::connect(listen).await.is_err());
}
#[test]
fn connection_diagnostics_do_not_render_peer_controlled_errors() {
let details = "private-peer-data";
let json_error = serde_json::from_value::<u64>(serde_json::json!(details))
.expect_err("invalid request field");
for (error, expected) in [
(Error::Protocol(details.into()), "protocol"),
(Error::Json(json_error), "JSON Data at 0:0"),
(
Error::Io(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
details,
)),
"I/O ConnectionReset",
),
] {
assert_eq!(connection_diagnostic(&error), expected);
}
}
#[test]
fn pre_auth_frames_reject_unknown_fields() {
let error = serde_json::from_value::<PreAuthClientFrame>(serde_json::json!({
"version": crate::wire::PROTOCOL_VERSION,
"type": "authenticate",
"token": "secret",
"client_kind": "cli",
"unexpected": true,
}))
.expect_err("reject unknown pre-auth field");
assert!(error.to_string().contains("unknown field `unexpected`"));
}
fn append_masked_binary_frame(output: &mut Vec<u8>, payload: &[u8]) {
output.push(0x82);
if payload.len() <= 125 {
output.push(0x80 | u8::try_from(payload.len()).expect("small WebSocket payload"));
} else {
output.push(0x80 | 126);
output.extend_from_slice(
&u16::try_from(payload.len())
.expect("WebSocket test payload fits u16")
.to_be_bytes(),
);
}
let mask = [0x12, 0x34, 0x56, 0x78];
output.extend_from_slice(&mask);
output.extend(
payload
.iter()
.enumerate()
.map(|(index, byte)| byte ^ mask[index % mask.len()]),
);
}
#[tokio::test]
async fn connection_admission_wakes_waiters_and_bounds_authenticated_clients() {
let admission = ConnectionAdmission::new(1, 1);
let first = admission.admit().await;
let waiting = admission.admit();
tokio::pin!(waiting);
tokio::select! {
biased;
_ = &mut waiting => panic!("second pre-auth connection bypassed the bound"),
() = std::future::ready(()) => {}
}
let authenticated = first.promote().expect("promote first connection");
let second = waiting.await;
assert!(second.promote().is_none());
drop(authenticated);
assert!(admission.admit().await.promote().is_some());
}
#[tokio::test]
async fn pre_auth_reader_rejects_an_oversized_frame_from_its_prefix() {
let (mut writer, reader) = tokio::io::duplex(4);
let mut reader = FrameReader::new(reader);
let oversized = u32::try_from(MAX_PRE_AUTH_FRAME_BYTES + 1).expect("frame limit fits u32");
writer
.write_all(&oversized.to_be_bytes())
.await
.expect("write prefix");
let error = read_frame_with_limit::<PreAuthClientFrame>(&mut reader, MAX_PRE_AUTH_FRAME_BYTES)
.await
.expect_err("oversized pre-auth frame must fail");
assert!(matches!(error, Error::Protocol(_)), "{error}");
}
#[test]
fn client_inventory_aggregates_connections_and_keeps_inactive_devices() {
let root = tempfile::tempdir().unwrap();
let bots = Arc::new(BotStore::open(root.path()).unwrap());
let clients = Arc::new(ClientConnections::default());
let identity = ClientIdentity {
id: "client-a".into(),
label: "Mac".into(),
};
let paired = [identity.clone()];
let first = clients
.register(identity.id.clone(), ClientKind::Macos, Arc::clone(&bots))
.expect("first connection");
let _dashboard = clients
.register(
identity.id.clone(),
ClientKind::GatewayDashboard,
Arc::clone(&bots),
)
.expect("dashboard connection");
let second = clients
.register(identity.id, ClientKind::Macos, bots)
.expect("second connection");
let two = clients.snapshot(&paired).expect("two connections")[0].connections;
assert_eq!(clients.native_count().unwrap(), 2);
drop(first);
let one = clients.snapshot(&paired).expect("one connection")[0].connections;
drop(second);
let inactive = clients.snapshot(&paired).expect("inactive client")[0].clone();
assert_eq!(
clients.native_count().unwrap(),
0,
"dashboard is not native activity"
);
assert_eq!(
(two, one, inactive.connections, inactive.kinds),
(2, 1, 0, Vec::new())
);
}
#[test]
fn client_presence_hooks_follow_first_and_last_native_socket_and_survive_restart() {
use crate::wire::{AgentComposition, HookData, HookSource, VersionedAgentConfig};
let root = tempfile::tempdir().unwrap();
let bots = Arc::new(BotStore::open(root.path()).unwrap());
let first_bot = bots
.seed_default(&VersionedAgentConfig {
revision: 1,
config: AgentComposition::default(),
})
.unwrap()
.unwrap();
let second_bot = bots
.create_bot(
"Helper",
"Own client lifecycle checks.",
AgentComposition::default(),
)
.unwrap();
let clients = Arc::new(ClientConnections::default());
let events = || bots.unpublished_events(100).unwrap();
let dashboard = clients
.register(
"client-a".into(),
ClientKind::GatewayDashboard,
Arc::clone(&bots),
)
.unwrap();
assert!(events().is_empty());
let first = clients
.register("client-a".into(), ClientKind::Macos, Arc::clone(&bots))
.unwrap();
assert_eq!(events().len(), 2);
let second = clients
.register("client-a".into(), ClientKind::Ios, Arc::clone(&bots))
.unwrap();
drop(first);
assert_eq!(
events().len(),
2,
"another native kind keeps the client connected"
);
drop(second);
let disconnected = events();
assert_eq!(disconnected.len(), 4);
for bot in [&first_bot, &second_bot] {
let owned: Vec<_> = disconnected
.iter()
.filter(|event| event.bot_id == bot.id)
.collect();
assert_eq!(owned.len(), 2);
assert!(matches!(owned[0].data, HookData::ClientConnected { .. }));
assert!(matches!(owned[1].data, HookData::ClientDisconnected { .. }));
assert!(owned.iter().all(|event| event.source
== HookSource::Client {
client_id: "client-a".into()
}
&& event.cause_id.is_none()
&& event.ancestry.is_empty()));
}
drop(dashboard);
assert_eq!(
events().len(),
4,
"dashboard disconnect is not native presence"
);
let reconnect = clients
.register("client-a".into(), ClientKind::Macos, Arc::clone(&bots))
.unwrap();
assert_eq!(events().len(), 6, "a later connection is a new transition");
drop(reconnect);
let ids: std::collections::BTreeSet<_> =
events().iter().map(|event| event.id.clone()).collect();
assert_eq!(ids.len(), 8);
let restarted = BotStore::open(root.path()).unwrap();
assert_eq!(restarted.unpublished_events(100).unwrap().len(), 8);
}
#[test]
fn websocket_upgrade_rejects_non_root_targets() {
let request = Request::builder().uri("/other").body(()).expect("request");
let rejection = WebSocketUpgradePolicy {
expected_host: None,
}
.on_request(&request, Response::new(()))
.expect_err("non-root path must fail");
assert_eq!(rejection.status(), StatusCode::NOT_FOUND);
}
#[test]
fn websocket_upgrade_rejects_browser_origins() {
let request = Request::builder()
.uri("/")
.header(ORIGIN, "https://attacker.example")
.body(())
.expect("request");
let rejection = WebSocketUpgradePolicy {
expected_host: None,
}
.on_request(&request, Response::new(()))
.expect_err("Origin header must fail");
assert_eq!(rejection.status(), StatusCode::FORBIDDEN);
}
#[test]
fn websocket_upgrade_rejects_the_wrong_cloudflare_host() {
let request = Request::builder()
.uri("/")
.header(HOST, "other.example")
.body(())
.expect("request");
let rejection = WebSocketUpgradePolicy {
expected_host: Some("gateway.example".into()),
}
.on_request(&request, Response::new(()))
.expect_err("wrong Host must fail");
assert_eq!(rejection.status(), StatusCode::FORBIDDEN);
}
#[test]
fn websocket_upgrade_accepts_the_cloudflare_host_with_standard_port() {
let request = Request::builder()
.uri("/")
.header(HOST, "gateway.example:443")
.header(SEC_WEBSOCKET_PROTOCOL, crate::channel::SUBPROTOCOL)
.body(())
.expect("request");
let accepted = WebSocketUpgradePolicy {
expected_host: Some("gateway.example".into()),
}
.on_request(&request, Response::new(()));
assert_eq!(
accepted.unwrap().headers()[SEC_WEBSOCKET_PROTOCOL],
crate::channel::SUBPROTOCOL
);
}
#[tokio::test]
async fn websocket_preserves_a_pipelined_bulk_frame_across_authentication() {
let root = tempfile::tempdir().expect("temporary directory");
let (server, grant) = GatewayServer::bootstrap(
root.path().join("state"),
std::net::SocketAddr::from(([127, 0, 0, 1], 0)),
)
.await
.expect("bootstrap gateway");
let listen = server.listen_addr();
let (shutdown, signal) = tokio::sync::oneshot::channel();
let serving = tokio::spawn(server.serve_until(async move {
let _ = signal.await;
}));
use tokio_tungstenite::tungstenite::client::IntoClientRequest as _;
let mut request = format!("ws://{listen}").into_client_request().unwrap();
request.headers_mut().insert(
SEC_WEBSOCKET_PROTOCOL,
crate::channel::SUBPROTOCOL.parse().unwrap(),
);
let (mut websocket, _) = tokio_tungstenite::connect_async(request).await.unwrap();
let mut state = crate::channel::client_handshake(&mut websocket, &grant.code)
.await
.unwrap();
let mut pairing = Vec::new();
write_frame(
&mut pairing,
&ClientFrame::new(ClientMessage::Pair {
code: grant.code,
client_label: "WebSocket test".into(),
client_kind: ClientKind::Ios,
}),
)
.await
.unwrap();
let request_id = "post-auth-list";
let mut post_auth = Vec::new();
write_frame(
&mut post_auth,
&ClientFrame::new(ClientMessage::UploadSessionFileChunk {
request_id: request_id.into(),
session_id: "missing-session".into(),
upload_id: "missing-upload".into(),
offset: 0,
data: vec![0; MAX_PRE_AUTH_FRAME_BYTES],
}),
)
.await
.unwrap();
assert!(post_auth.len() > MAX_PRE_AUTH_FRAME_BYTES);
let mut pipelined = Vec::new();
for plaintext in [pairing, post_auth] {
let mut encrypted = vec![0; plaintext.len() + crate::channel::TAG_BYTES];
let length = state.write_message(&plaintext, &mut encrypted).unwrap();
append_masked_binary_frame(&mut pipelined, &encrypted[..length]);
}
websocket
.get_mut()
.write_all(&pipelined)
.await
.expect("pipeline encrypted pairing and post-auth records");
let (transport, stream) = tokio::io::duplex(16 * 1024);
let bridge = tokio::spawn(crate::channel::bridge(websocket, state, stream));
let mut reader = FrameReader::new(transport);
let paired = read_frame::<ServerFrame>(&mut reader)
.await
.unwrap()
.unwrap();
let authenticated = read_frame::<ServerFrame>(&mut reader)
.await
.unwrap()
.unwrap();
let ready = read_frame::<ServerFrame>(&mut reader)
.await
.unwrap()
.unwrap();
let rejection = read_frame::<ServerFrame>(&mut reader)
.await
.unwrap()
.unwrap();
assert!(matches!(
(
paired.message,
authenticated.message,
ready.message,
rejection.message
),
(
ServerMessage::Paired { .. },
ServerMessage::Authenticated,
ServerMessage::Ready { .. },
ServerMessage::Rejected { request_id: actual, .. }
) if actual == request_id
));
bridge.abort();
shutdown.send(()).expect("stop gateway");
serving.await.expect("gateway task").expect("gateway stop");
}
#[tokio::test(start_paused = true)]
async fn websocket_upgrade_and_authentication_share_one_deadline() {
let root = tempfile::tempdir().expect("temporary directory");
let (server, _) = GatewayServer::bootstrap(
root.path().join("state"),
std::net::SocketAddr::from(([127, 0, 0, 1], 0)),
)
.await
.expect("bootstrap gateway");
let listen = server.listen_addr();
let GatewayServer {
listener,
auth,
host,
bots,
..
} = server;
let client_connections = Arc::new(ClientConnections::default());
let (client_revocations, _) = broadcast::channel(MAX_CONNECTIONS);
let admission = ConnectionAdmission::new(1, 1).admit().await;
let (accepted_tx, accepted_rx) = tokio::sync::oneshot::channel();
let serving = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept connection");
let auth_deadline = Instant::now() + PRE_AUTH_TIMEOUT;
accepted_tx.send(()).expect("report accepted connection");
serve_plaintext_connection(
stream,
ConnectionContext {
local: true,
auth,
host,
bots,
client_connections,
client_revocations,
admission,
access_lease: None,
},
PlaintextHandshake {
expected_websocket_host: None,
auth_deadline,
},
)
.await
});
let mut stream = TcpStream::connect(listen).await.expect("connect gateway");
accepted_rx.await.expect("connection accepted");
tokio::time::advance(Duration::from_secs(2)).await;
stream.write_all(b"G").await.expect("start upgrade");
tokio::task::yield_now().await;
tokio::time::advance(Duration::from_secs(2)).await;
stream
.write_all(
format!(
"ET / HTTP/1.1\r\nHost: {listen}\r\nConnection: Upgrade\r\nUpgrade: websocket\r\nSec-WebSocket-Version: 13\r\nSec-WebSocket-Protocol: mobius-noise-v1\r\nSec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\r\n"
)
.as_bytes(),
)
.await
.expect("finish upgrade");
let mut response = Vec::new();
while !response.ends_with(b"\r\n\r\n") {
let mut byte = [0_u8; 1];
let read = stream.read(&mut byte).await.expect("read upgrade response");
assert_eq!(read, 1, "upgrade response ended early");
response.push(byte[0]);
}
assert!(response.starts_with(b"HTTP/1.1 101"));
let mut websocket = WebSocketStream::from_raw_socket(stream, Role::Client, None).await;
tokio::time::advance(Duration::from_secs(2)).await;
tokio::time::resume();
let closed = tokio::time::timeout(Duration::from_secs(1), websocket.next())
.await
.expect("authentication deadline must close the socket");
assert!(matches!(
closed,
None | Some(Ok(Message::Close(_))) | Some(Err(_))
));
assert!(matches!(
serving.await.expect("gateway task"),
Err(Error::Unauthorized)
));
}
#[tokio::test]
async fn bootstrap_owns_the_listener_before_creating_state() {
let root = tempfile::tempdir().expect("temporary directory");
let occupied = TcpListener::bind("127.0.0.1:0")
.await
.expect("occupied listener");
let listen = occupied.local_addr().expect("listen address");
let state = root.path().join("state");
let result = GatewayServer::bootstrap(state.clone(), listen).await;
assert!(matches!(result, Err(Error::Io(_))));
assert!(!state.exists());
}
#[tokio::test]
async fn connected_client_pauses_and_resets_inactivity_shutdown() {
let root = tempfile::tempdir().expect("temporary directory");
let (server, grant) = GatewayServer::bootstrap(
root.path().join("state"),
std::net::SocketAddr::from(([127, 0, 0, 1], 0)),
)
.await
.expect("bootstrap gateway");
let listen = server.config.listen;
let serving = tokio::spawn(
server.serve_until_inactive(std::future::pending(), Duration::from_millis(200)),
);
let endpoint = format!("tcp://{listen}")
.parse::<Endpoint>()
.expect("endpoint");
let (connection, _) =
GatewayClient::pair(&endpoint, grant.code, "inactivity test", ClientKind::Cli)
.await
.expect("connect client");
tokio::time::sleep(Duration::from_millis(300)).await;
assert!(!serving.is_finished());
drop(connection);
tokio::time::sleep(Duration::from_millis(75)).await;
assert!(!serving.is_finished());
tokio::time::timeout(Duration::from_secs(2), serving)
.await
.expect("inactivity shutdown timeout")
.expect("gateway task")
.expect("gateway shutdown");
}