use std::{sync::Arc, time::Duration};
use base64::{Engine, engine::general_purpose::STANDARD};
use keydous_bridge::{
catalog::SimulatedCatalog,
driver::{DeviceList, dj_dev::Oneofdev},
server::{ServerConfig, serve},
};
use prost::Message;
use tokio::{net::TcpListener, sync::oneshot};
const OFFICIAL_ORIGIN: &str = "https://keydousnj.rongyuan.tech";
#[tokio::test]
async fn official_origin_receives_discovery_stream() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let config = ServerConfig {
address,
allowed_origins: vec![OFFICIAL_ORIGIN.into()],
};
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let server = tokio::spawn(serve(
listener,
config,
Arc::new(SimulatedCatalog),
async move {
let _ = shutdown_rx.await;
},
));
let response = reqwest::Client::new()
.post(format!("http://{address}/driver.DriverGrpc/watchDevList"))
.header("origin", OFFICIAL_ORIGIN)
.header("content-type", "application/grpc-web-text")
.header("accept", "application/grpc-web-text")
.header("x-grpc-web", "1")
.body("AAAAAAA=")
.send()
.await
.unwrap();
assert_eq!(response.status(), 200);
assert_eq!(
response
.headers()
.get("access-control-allow-origin")
.unwrap(),
OFFICIAL_ORIGIN
);
let body = tokio::time::timeout(Duration::from_secs(2), response.text())
.await
.unwrap()
.unwrap();
let frames = decode_frames(&body);
assert_eq!(frames[0].0, 0);
let snapshot = DeviceList::decode(frames[0].1.as_slice()).unwrap();
let Some(Oneofdev::Dev(device)) = &snapshot.devlist[0].oneofdev else {
panic!("expected direct device");
};
assert_eq!((device.vid, device.pid), (0x3151, 0x5030));
assert_eq!(frames.last().unwrap().0, 0x80);
assert_eq!(
std::str::from_utf8(&frames.last().unwrap().1).unwrap(),
"grpc-status:0\r\n"
);
let _ = shutdown_tx.send(());
server.await.unwrap().unwrap();
}
#[tokio::test]
async fn non_loopback_listener_is_rejected_even_with_loopback_config() {
let listener = TcpListener::bind("0.0.0.0:0").await.unwrap();
let config = ServerConfig {
address: "127.0.0.1:0".parse().unwrap(),
allowed_origins: vec![OFFICIAL_ORIGIN.into()],
};
let server = tokio::spawn(serve(
listener,
config,
Arc::new(SimulatedCatalog),
std::future::pending(),
));
let result = tokio::time::timeout(Duration::from_millis(200), server)
.await
.expect("non-loopback listener was not rejected")
.expect_err("non-loopback listener served without panicking");
assert!(result.is_panic());
}
#[tokio::test]
async fn listener_address_must_match_nonzero_config_address() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let listener_port = listener.local_addr().unwrap().port();
let mismatched_port = if listener_port == u16::MAX {
listener_port - 1
} else {
listener_port + 1
};
let config = ServerConfig {
address: format!("127.0.0.1:{mismatched_port}").parse().unwrap(),
allowed_origins: vec![OFFICIAL_ORIGIN.into()],
};
let server = tokio::spawn(serve(
listener,
config,
Arc::new(SimulatedCatalog),
std::future::pending(),
));
let result = tokio::time::timeout(Duration::from_millis(200), server)
.await
.expect("mismatched listener address was not rejected")
.expect_err("mismatched listener address served without panicking");
assert!(result.is_panic());
}
#[tokio::test]
async fn zero_config_port_does_not_wildcard_listener_port() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let config = ServerConfig {
address: "127.0.0.1:0".parse().unwrap(),
allowed_origins: vec![OFFICIAL_ORIGIN.into()],
};
let server = tokio::spawn(serve(
listener,
config,
Arc::new(SimulatedCatalog),
std::future::pending(),
));
let result = tokio::time::timeout(Duration::from_millis(200), server)
.await
.expect("zero configured port was accepted as a wildcard")
.expect_err("zero configured port served without panicking");
assert!(result.is_panic());
}
#[tokio::test]
async fn disallowed_origin_is_denied_by_browser_cors_enforcement() {
let (address, shutdown_tx, server) = start_server().await;
let response = reqwest::Client::new()
.post(format!("http://{address}/driver.DriverGrpc/watchDevList"))
.header("origin", "https://attacker.invalid")
.header("content-type", "application/grpc-web-text")
.header("x-grpc-web", "1")
.body("AAAAAAA=")
.send()
.await
.unwrap();
assert_eq!(response.status(), 200);
assert!(
response
.headers()
.get("access-control-allow-origin")
.is_none()
);
let _ = shutdown_tx.send(());
server.await.unwrap().unwrap();
}
#[tokio::test]
async fn official_origin_private_network_preflight_is_allowed() {
let (address, shutdown_tx, server) = start_server().await;
let response = reqwest::Client::new()
.request(
reqwest::Method::OPTIONS,
format!("http://{address}/driver.DriverGrpc/watchDevList"),
)
.header("origin", OFFICIAL_ORIGIN)
.header("access-control-request-method", "POST")
.header(
"access-control-request-headers",
"content-type,x-grpc-web,x-user-agent",
)
.header("access-control-request-private-network", "true")
.send()
.await
.unwrap();
assert_eq!(response.status(), 200);
assert_eq!(
response
.headers()
.get("access-control-allow-private-network")
.unwrap(),
"true"
);
assert_eq!(
response
.headers()
.get("access-control-allow-origin")
.unwrap(),
OFFICIAL_ORIGIN
);
assert!(
response
.headers()
.get("access-control-allow-methods")
.unwrap()
.to_str()
.unwrap()
.split(',')
.any(|method| method.trim() == "POST")
);
let allowed_headers = response
.headers()
.get("access-control-allow-headers")
.unwrap()
.to_str()
.unwrap();
assert!(allowed_headers.contains("content-type"));
assert!(allowed_headers.contains("x-grpc-web"));
assert!(allowed_headers.contains("x-user-agent"));
let _ = shutdown_tx.send(());
server.await.unwrap().unwrap();
}
async fn start_server() -> (
std::net::SocketAddr,
oneshot::Sender<()>,
tokio::task::JoinHandle<Result<(), tonic::transport::Error>>,
) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let config = ServerConfig {
address,
allowed_origins: vec![OFFICIAL_ORIGIN.into()],
};
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let server = tokio::spawn(serve(
listener,
config,
Arc::new(SimulatedCatalog),
async move {
let _ = shutdown_rx.await;
},
));
(address, shutdown_tx, server)
}
fn decode_frames(body: &str) -> Vec<(u8, Vec<u8>)> {
let mut bytes = Vec::new();
let mut remaining = body;
while !remaining.is_empty() {
let chunk_end = remaining
.find('=')
.map(|index| {
index
+ remaining[index..]
.bytes()
.take_while(|byte| *byte == b'=')
.count()
})
.unwrap_or(remaining.len());
bytes.extend(STANDARD.decode(&remaining[..chunk_end]).unwrap());
remaining = &remaining[chunk_end..];
}
let mut frames = Vec::new();
let mut offset = 0;
while offset < bytes.len() {
let flag = bytes[offset];
let length = u32::from_be_bytes(bytes[offset + 1..offset + 5].try_into().unwrap()) as usize;
let payload_start = offset + 5;
let payload_end = payload_start + length;
frames.push((flag, bytes[payload_start..payload_end].to_vec()));
offset = payload_end;
}
frames
}