keydous-bridge 0.1.2

Linux bridge for configuring Keydous keyboards with the official web driver
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
}