breakmancer 0.9.0

Drop a breakpoint into any shell.
Documentation
//! Integration test exercising a connection setup, handshake, single
//! successful command run, and shutdown. (Does not check if transport
//! connection gets closed, though.)

use std::time::Duration;

use tokio::net::TcpListener;

use breakmancer::{
    protocol::{
        breakpoint_open_connection, controller_open_connection, BreakpointId, Command, Connection,
        ControllerId, ExitBreak, Finished, IdString, OutputLine, OutputType, SecureMsg,
        XWingKeypair,
    },
    transport::{
        tcp::{TcpClient, TcpServer},
        wormhole::{WormholeCaller, WormholeListener},
        TransportCaller, TransportListener,
    },
};

async fn setup_controller(
    ctrl_keys: XWingKeypair,
    break_digest: BreakpointId,
    mut listener: TransportListener,
) -> Result<Connection, String> {
    // Blocks on breakpoint connecting.
    let half_ctrl = controller_open_connection(ctrl_keys.clone(), &mut listener).await?;
    let (conn, break_intro) = half_ctrl
        .verify_breakpoint_and_complete_setup(&break_digest)
        .await?;

    assert_eq!(break_intro.which, Some("Integration test".to_string()));

    Ok(conn)
}

async fn exercise_controller(conn: &mut Connection) {
    conn.send(SecureMsg::Command(Command {
        command: "echo $(( 1 + 2 ))".to_string(),
        seq: 0,
    }))
    .await
    .unwrap();

    match conn.recv().await.unwrap() {
        Some(SecureMsg::OutputLine(out)) => {
            assert_eq!(out.seq, 0);
            assert_eq!(out.bytes, "3".as_bytes());
            assert_eq!(out.gender, OutputType::Stdout);
        }
        Some(wrong) => panic!("Wrong message when expecting output: {wrong:?}"),
        None => panic!("Connection closed before output received."),
    }

    match conn.recv().await.unwrap() {
        Some(SecureMsg::Finished(out)) => {
            assert_eq!(out.seq, 0);
            assert_eq!(out.exit_code, Some(0));
        }
        Some(wrong) => panic!("Wrong message when expecting command-finish: {wrong:?}"),
        None => panic!("Connection closed before command-finish received."),
    }

    conn.send(SecureMsg::ExitBreak(ExitBreak {})).await.unwrap();
}

async fn setup_breakpoint(
    break_keys: XWingKeypair,
    ctrl_digest: ControllerId,
    break_digest: BreakpointId,
    caller: TransportCaller,
) -> Result<Connection, String> {
    let half_break = breakpoint_open_connection(
        &break_keys,
        &caller,
        &ctrl_digest,
        Some("Integration test".to_string()),
    )
    .await
    .map_err(|err| format!("Couldn't open connection: {err:?}"))?;

    // Ensure both sides are producing/expecting the same thing.
    // NOTE: In real code we'd want to use constant-time compare
    // for checking equality of digests (`ids_match` function.)
    assert_eq!(
        half_break.breakpoint_id.id_string(),
        break_digest.id_string()
    );

    half_break
        .continue_after_id_printed()
        .await
        .map_err(|err| format!("Couldn't finish handshake: {err:?}"))
}

async fn exercise_breakpoint(conn: &mut Connection) {
    match conn.recv().await.unwrap() {
        Some(SecureMsg::Command(cmd)) => {
            assert_eq!(cmd.seq, 0);
            assert_eq!(cmd.command, "echo $(( 1 + 2 ))".to_string());
        }
        Some(wrong) => panic!("Wrong message when expecting command: {wrong:?}"),
        None => panic!("Connection closed before command received."),
    }

    conn.send(SecureMsg::OutputLine(OutputLine {
        seq: 0,
        gender: OutputType::Stdout,
        bytes: "3".as_bytes().to_vec(),
    }))
    .await
    .unwrap();

    conn.send(SecureMsg::Finished(Finished {
        seq: 0,
        exit_code: Some(0),
    }))
    .await
    .unwrap();

    match conn.recv().await.unwrap() {
        Some(SecureMsg::ExitBreak(_)) => (),
        Some(wrong) => panic!("Wrong message when expecting exitbreak: {wrong:?}"),
        None => panic!("Connection closed before exitbreak received."),
    }
}

async fn test_generic_happy_path(
    ctrl_keys: XWingKeypair,
    break_keys: XWingKeypair,
    listener: TransportListener,
    caller: TransportCaller,
) -> Result<(), String> {
    let ctrl_digest = ControllerId::from_key(&ctrl_keys.public_key);
    let break_digest_1 = BreakpointId::from_key(&break_keys.public_key);
    let break_digest_2 = break_digest_1.clone();

    let ctrl_proc = tokio::spawn(async move {
        let mut conn = setup_controller(ctrl_keys, break_digest_1, listener)
            .await
            .unwrap();
        exercise_controller(&mut conn).await
    });
    let break_proc = tokio::spawn(async move {
        let mut conn = setup_breakpoint(break_keys, ctrl_digest, break_digest_2, caller)
            .await
            .unwrap();
        exercise_breakpoint(&mut conn).await
    });

    tokio::select! {
        result = ctrl_proc => {
            result.map_err(|err| format!("Error in controller process: {err}"))?
        }

        _ = tokio::time::sleep(Duration::from_secs(15)) => {
            Err("Timed out waiting for controller to finish".to_string())?
        }
    };

    tokio::select! {
        result = break_proc => {
            result.map_err(|err| format!("Error in breakpoint process: {err}"))
        }

        _ = tokio::time::sleep(Duration::from_secs(15)) => {
            Err("Timed out waiting for breakpoint to finish".to_string())
        }
    }?;

    Ok(())
}

#[tokio::test]
async fn generic_happy_path_tcp() -> Result<(), String> {
    let ctrl_keys = XWingKeypair::new();
    let break_keys = XWingKeypair::new();

    let server_socket = TcpListener::bind("[::]:0")
        .await
        .map_err(|err| format!("Could not open a port to listen on: {err}"))?;
    let local_port = server_socket.local_addr().unwrap().port();

    let listener = TransportListener::Tcp(TcpServer::from_socket(server_socket));
    let caller = TransportCaller::Tcp(TcpClient::new(&format!("localhost:{local_port}")));

    test_generic_happy_path(ctrl_keys, break_keys, listener, caller).await
}

#[tokio::test]
async fn generic_happy_path_wormhole() -> Result<(), String> {
    let ctrl_keys = XWingKeypair::new();
    let break_keys = XWingKeypair::new();

    let controller_id = ControllerId::from_key(&ctrl_keys.public_key);
    let listener = TransportListener::Wormhole(WormholeListener::new(&controller_id));
    let caller = TransportCaller::Wormhole(WormholeCaller::new(&controller_id));

    test_generic_happy_path(ctrl_keys, break_keys, listener, caller).await
}