keydous-bridge 0.1.2

Linux bridge for configuring Keydous keyboards with the official web driver
use std::{
    future::Future,
    net::{IpAddr, Ipv4Addr, SocketAddr},
    sync::Arc,
};

use tokio::net::TcpListener;
use tokio_stream::wrappers::TcpListenerStream;
use tonic::{codegen::http::Method, transport::Server};
use tower_http::cors::{AllowOrigin, CorsLayer};

use crate::{
    catalog::DeviceCatalog,
    driver::driver_grpc_server::DriverGrpcServer,
    service::DriverService,
    transport::{DeniedDeviceIo, DeviceIo},
};

#[derive(Clone, Debug, Eq, PartialEq)]
pub struct ServerConfig {
    pub address: SocketAddr,
    pub allowed_origins: Vec<String>,
}

impl ServerConfig {
    pub fn official() -> Self {
        Self {
            address: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 3814),
            allowed_origins: vec!["https://keydousnj.rongyuan.tech".into()],
        }
    }
}

pub async fn serve(
    listener: TcpListener,
    config: ServerConfig,
    catalog: Arc<dyn DeviceCatalog>,
    shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), tonic::transport::Error> {
    serve_with_device_io(
        listener,
        config,
        catalog,
        Arc::new(DeniedDeviceIo),
        shutdown,
    )
    .await
}

pub async fn serve_with_device_io(
    listener: TcpListener,
    config: ServerConfig,
    catalog: Arc<dyn DeviceCatalog>,
    device_io: Arc<dyn DeviceIo>,
    shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), tonic::transport::Error> {
    let listener_address = listener
        .local_addr()
        .expect("listener must have a local address");
    assert!(
        config.address.ip().is_loopback(),
        "server config must be loopback-only"
    );
    assert!(
        listener_address.ip().is_loopback(),
        "server listener must be loopback-only"
    );
    assert_eq!(
        listener_address, config.address,
        "server listener address must match config"
    );
    let origins = config
        .allowed_origins
        .iter()
        .map(|origin| origin.parse().expect("valid configured origin"))
        .collect::<Vec<_>>();
    let cors = CorsLayer::new()
        .allow_origin(AllowOrigin::list(origins))
        .allow_methods([Method::POST])
        .allow_headers([
            "content-type".parse().unwrap(),
            "x-grpc-web".parse().unwrap(),
            "x-user-agent".parse().unwrap(),
            "grpc-timeout".parse().unwrap(),
            "access-control-request-private-network".parse().unwrap(),
        ])
        .expose_headers([
            "grpc-status".parse().unwrap(),
            "grpc-message".parse().unwrap(),
        ])
        .allow_private_network(true);
    let service = DriverGrpcServer::new(DriverService::with_device_io(catalog, device_io));

    Server::builder()
        .accept_http1(true)
        .layer(cors)
        .layer(tonic_web::GrpcWebLayer::new())
        .add_service(service)
        .serve_with_incoming_shutdown(TcpListenerStream::new(listener), shutdown)
        .await
}

#[cfg(test)]
mod tests {
    use super::ServerConfig;

    #[test]
    fn official_config_is_loopback_only() {
        let config = ServerConfig::official();
        assert!(config.address.ip().is_loopback());
        assert_eq!(config.address.port(), 3814);
        assert_eq!(config.allowed_origins, ["https://keydousnj.rongyuan.tech"]);
    }
}