#[path = "../../src/proto/codex.exec_server.relay.v1.rs"] mod relay_proto; use std::sync::Arc; use std::sync::Mutex; use std::time::Duration; use anyhow::Context; use anyhow::Result; use codex_exec_server::NoiseChannelPublicKey; use futures::SinkExt; use futures::StreamExt; use prost::Message as ProstMessage; use relay_proto::RelayMessageFrame; use relay_proto::relay_message_frame; use tokio::net::TcpListener; use tokio::net::TcpStream; use tokio::time::timeout; use tokio_tungstenite::WebSocketStream; use tokio_tungstenite::accept_async; use tokio_tungstenite::tungstenite::Message; use wiremock::MockServer; pub const TEST_TIMEOUT: Duration = Duration::from_secs(30); pub async fn accept_websocket( listener: &TcpListener, role: &str, ) -> Result> { let (socket, _peer_addr) = timeout(TEST_TIMEOUT, listener.accept()) .await .with_context(|| format!("remote {role} should connect to fake rendezvous"))??; timeout(TEST_TIMEOUT, accept_async(socket)) .await .with_context(|| format!("fake rendezvous should accept {role} websocket"))? .map_err(Into::into) } pub async fn registered_executor_public_key( registry: &MockServer, ) -> Result { let requests = registry .received_requests() .await .context("wiremock should retain requests")?; let request = requests .iter() .find(|request| request.url.path().ends_with("/register")) .context("exec-server should register before connecting")?; let body: serde_json::Value = serde_json::from_slice(&request.body)?; let key = serde_json::from_value(body["executor_public_key"].clone())?; Ok(key) } pub async fn proxy_relay_frames( mut environment: WebSocketStream, mut harness: WebSocketStream, captured_frames: Arc>>>, ) -> Result<()> { loop { tokio::select! { message = environment.next() => { let Some(message) = message else { break; }; let message = message?; capture_binary_frame(&captured_frames, &message); harness.send(message).await?; } message = harness.next() => { let Some(message) = message else { break; }; let message = message?; capture_binary_frame(&captured_frames, &message); environment.send(message).await?; } } } Ok(()) } fn capture_binary_frame(captured_frames: &Mutex>>, message: &Message) { if let Message::Binary(bytes) = message { captured_frames .lock() .unwrap_or_else(std::sync::PoisonError::into_inner) .push(bytes.to_vec()); } } pub fn assert_relay_data_is_encrypted(captured_frames: &Mutex>>) -> Result<()> { let captured_frames = captured_frames .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); let mut data_frames = 0; for encoded in captured_frames.iter() { let frame = RelayMessageFrame::decode(encoded.as_slice())?; let Some(relay_message_frame::Body::Data(data)) = frame.body else { continue; }; data_frames += 1; let payload = String::from_utf8_lossy(&data.payload); assert!(!payload.contains("initialize")); assert!(!payload.contains("process/start")); assert!(!payload.contains("noise-relay-test")); } assert!( data_frames >= 4, "expected encrypted request and response frames" ); Ok(()) }