mirror of
https://github.com/openai/codex.git
synced 2026-09-06 15:29:32 +00:00
## Why Allow remote exec servers to connect directly to AWS-hosted registries that authenticate registry requests and WebSocket handshakes with AWS SigV4. ## What changed - Add `--remote-transport direct` with SigV4 profile, region, and service options while keeping Noise as the default transport. - Register the `direct_jsonrpc_v1` transport and carry plain exec-server JSON-RPC messages over the authenticated WebSocket. - Reuse direct registrations across transient disconnects, refresh them after a `409 Conflict`, and require TLS for non-loopback endpoints. ## Testing - Cover CLI validation and SigV4 request signing. - Exercise direct registration, handshake retry behavior, JSON-RPC interoperability, and process recovery after reconnecting. GitOrigin-RevId: 0755df330ba3abe5db0a516fdaa49338d9bbe2d2
807 lines
28 KiB
Rust
807 lines
28 KiB
Rust
mod common;
|
|
|
|
use std::collections::HashMap;
|
|
use std::sync::Arc;
|
|
use std::time::Duration;
|
|
|
|
use anyhow::Context;
|
|
use anyhow::Result;
|
|
use axum::Router;
|
|
use axum::extract::State;
|
|
use axum::extract::WebSocketUpgrade;
|
|
use axum::response::IntoResponse;
|
|
use axum::routing::any;
|
|
use codex_api::AuthProvider;
|
|
#[cfg(unix)]
|
|
use codex_exec_server::EnvironmentConnectionState;
|
|
use codex_exec_server::EnvironmentInfo;
|
|
use codex_exec_server::EnvironmentManager;
|
|
use codex_exec_server::EnvironmentObservedStatus;
|
|
use codex_exec_server::EnvironmentStatus;
|
|
use codex_exec_server::EnvironmentStatusKind;
|
|
use codex_exec_server::ExecParams;
|
|
#[cfg(unix)]
|
|
use codex_exec_server::ExecProcessEvent;
|
|
use codex_exec_server::ExecResponse;
|
|
use codex_exec_server::ExecServerClientConnectOptions;
|
|
use codex_exec_server::ExecServerRuntimePaths;
|
|
use codex_exec_server::InitializeParams;
|
|
use codex_exec_server::InitializeResponse;
|
|
use codex_exec_server::ProcessId;
|
|
use codex_exec_server::ReadParams;
|
|
use codex_exec_server::ReadResponse;
|
|
use codex_exec_server::RemoteEnvironmentConfig;
|
|
use codex_exec_server::RemoteEnvironmentTransport;
|
|
#[cfg(unix)]
|
|
use codex_exec_server::WriteStatus;
|
|
use codex_exec_server_protocol::JSONRPCError;
|
|
use codex_exec_server_protocol::JSONRPCErrorError;
|
|
use codex_exec_server_protocol::JSONRPCMessage;
|
|
use codex_exec_server_protocol::JSONRPCNotification;
|
|
use codex_exec_server_protocol::JSONRPCRequest;
|
|
use codex_exec_server_protocol::JSONRPCResponse;
|
|
use codex_http_client::HttpClientFactory;
|
|
use codex_http_client::OutboundProxyPolicy;
|
|
use codex_utils_path_uri::PathUri;
|
|
use common::exec_server::DisconnectableWebSocketProxy;
|
|
use futures::SinkExt;
|
|
use futures::StreamExt;
|
|
use http::HeaderMap;
|
|
use http::HeaderValue;
|
|
use pretty_assertions::assert_eq;
|
|
use tokio::net::TcpListener;
|
|
use tokio::sync::mpsc;
|
|
use tokio::task::JoinHandle;
|
|
use tokio::time::timeout;
|
|
use tokio_tungstenite::MaybeTlsStream;
|
|
use tokio_tungstenite::WebSocketStream;
|
|
use tokio_tungstenite::connect_async;
|
|
use tokio_tungstenite::tungstenite::Message;
|
|
use tokio_util::task::AbortOnDropHandle;
|
|
use wiremock::Mock;
|
|
use wiremock::MockServer;
|
|
use wiremock::ResponseTemplate;
|
|
use wiremock::matchers::header;
|
|
use wiremock::matchers::method;
|
|
use wiremock::matchers::path;
|
|
|
|
const TEST_TIMEOUT: Duration = Duration::from_secs(5);
|
|
|
|
type AcceptedSocket = axum::extract::ws::WebSocket;
|
|
const SESSION_ALREADY_ATTACHED_ERROR_CODE: i64 = -32010;
|
|
|
|
#[derive(Debug)]
|
|
struct DirectExecutorAuth;
|
|
|
|
impl AuthProvider for DirectExecutorAuth {
|
|
fn add_auth_headers(&self, headers: &mut HeaderMap) {
|
|
headers.insert(
|
|
http::header::AUTHORIZATION,
|
|
HeaderValue::from_static("AWS4-HMAC-SHA256 test-signature"),
|
|
);
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn accepted_websocket_rejects_initial_resume_session_id() -> Result<()> {
|
|
let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?;
|
|
let (_socket, _) = connect_async(&websocket_url).await?;
|
|
let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv())
|
|
.await?
|
|
.context("accepted websocket channel should remain open")?;
|
|
let mut options = accepted_options();
|
|
options.resume_session_id = Some("session-1".to_string());
|
|
|
|
let error = EnvironmentManager::from_accepted_websocket(
|
|
"environment-1".to_string(),
|
|
accepted_websocket,
|
|
options,
|
|
HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault),
|
|
)
|
|
.await
|
|
.expect_err("initial accepted websocket should reject a resume session ID");
|
|
|
|
assert!(
|
|
error
|
|
.to_string()
|
|
.contains("initial connection cannot resume a session"),
|
|
"unexpected error: {error}"
|
|
);
|
|
|
|
server_task.abort();
|
|
let _ = server_task.await;
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn accepted_websocket_environment_info_uses_initialization_metadata() -> Result<()> {
|
|
let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?;
|
|
let (_socket, manager) =
|
|
connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?;
|
|
let environment = manager
|
|
.default_environment()
|
|
.context("accepted environment")?;
|
|
|
|
assert_eq!(
|
|
timeout(TEST_TIMEOUT, environment.info()).await??,
|
|
EnvironmentInfo::local()
|
|
);
|
|
|
|
server_task.abort();
|
|
let _ = server_task.await;
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn accepted_websocket_interoperates_and_recovers_with_real_direct_executor() -> Result<()> {
|
|
let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?;
|
|
let proxy = DisconnectableWebSocketProxy::new(&websocket_url).await?;
|
|
let registry = MockServer::start().await;
|
|
Mock::given(method("POST"))
|
|
.and(path("/cloud/environment/environment-1/direct/register"))
|
|
.and(header("authorization", "AWS4-HMAC-SHA256 test-signature"))
|
|
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
|
|
"environment_id": "environment-1",
|
|
"transport": "direct_jsonrpc_v1",
|
|
"registration_id": "registration-1",
|
|
"url": proxy.websocket_url(),
|
|
})))
|
|
.expect(1)
|
|
.mount(®istry)
|
|
.await;
|
|
|
|
let http_client_factory = HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault);
|
|
let config = RemoteEnvironmentConfig::new_with_transport(
|
|
registry.uri(),
|
|
"environment-1".to_string(),
|
|
RemoteEnvironmentTransport::Direct,
|
|
Arc::new(DirectExecutorAuth),
|
|
http_client_factory.clone(),
|
|
)?;
|
|
let (codex_exe, codex_linux_sandbox_exe) = common::current_test_binary_helper_paths()?;
|
|
let runtime_paths = ExecServerRuntimePaths::new(codex_exe, codex_linux_sandbox_exe)?;
|
|
let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel();
|
|
let executor_task = AbortOnDropHandle::new(tokio::spawn(
|
|
codex_exec_server::run_remote_environment_until_shutdown(
|
|
config,
|
|
runtime_paths,
|
|
async move {
|
|
let _ = shutdown_rx.await;
|
|
},
|
|
),
|
|
));
|
|
let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv())
|
|
.await?
|
|
.context("direct executor websocket should be accepted")?;
|
|
let manager = timeout(
|
|
TEST_TIMEOUT,
|
|
EnvironmentManager::from_accepted_websocket(
|
|
"environment-1".to_string(),
|
|
accepted_websocket,
|
|
accepted_options(),
|
|
http_client_factory,
|
|
),
|
|
)
|
|
.await??;
|
|
let environment = manager
|
|
.default_environment()
|
|
.context("direct executor environment should be installed")?;
|
|
|
|
assert_eq!(
|
|
timeout(TEST_TIMEOUT, environment.force_info()).await??,
|
|
EnvironmentInfo::local()
|
|
);
|
|
let files = tempfile::tempdir()?;
|
|
let large_file_path = files.path().join("large-response.bin");
|
|
let large_file_contents = vec![0x5a; 128 * 1024];
|
|
tokio::fs::write(&large_file_path, &large_file_contents).await?;
|
|
assert_eq!(
|
|
timeout(
|
|
TEST_TIMEOUT,
|
|
environment.get_filesystem().read_file(
|
|
&PathUri::from_host_native_path(&large_file_path)?,
|
|
Default::default(),
|
|
/*sandbox*/ None,
|
|
)
|
|
)
|
|
.await??,
|
|
large_file_contents,
|
|
);
|
|
|
|
// The process fixture uses a POSIX shell; metadata and shutdown remain tested on all platforms.
|
|
#[cfg(unix)]
|
|
{
|
|
let mut proxy = proxy;
|
|
let backend = environment.get_exec_backend();
|
|
let temp_dir = tempfile::TempDir::new()?;
|
|
let gate_path = temp_dir.path().join("release-output");
|
|
let emitted_path = temp_dir.path().join("output-emitted");
|
|
let session = timeout(
|
|
TEST_TIMEOUT,
|
|
backend.start(ExecParams {
|
|
metadata: Default::default(),
|
|
process_id: ProcessId::from("proc-recover"),
|
|
argv: vec![
|
|
"/bin/sh".to_string(),
|
|
"-c".to_string(),
|
|
concat!(
|
|
"printf 'ready:%s\\n' \"$$\"; ",
|
|
"while [ ! -f \"$GATE\" ]; do /bin/sleep 0.01; done; ",
|
|
"printf 'during:%s\\n' \"$$\"; ",
|
|
": > \"$EMITTED\"; ",
|
|
"IFS= read -r line; ",
|
|
"printf 'after:%s:%s\\n' \"$$\" \"$line\"; ",
|
|
"exit 7",
|
|
)
|
|
.to_string(),
|
|
],
|
|
cwd: PathUri::from_host_native_path(std::env::current_dir()?)?,
|
|
shell_snapshot: None,
|
|
env_policy: /*env_policy*/ None,
|
|
env: HashMap::from([
|
|
(
|
|
"GATE".to_string(),
|
|
gate_path.to_string_lossy().into_owned(),
|
|
),
|
|
(
|
|
"EMITTED".to_string(),
|
|
emitted_path.to_string_lossy().into_owned(),
|
|
),
|
|
]),
|
|
tty: false,
|
|
pipe_stdin: true,
|
|
arg0: None,
|
|
sandbox: None,
|
|
enforce_managed_network: false,
|
|
managed_network: None,
|
|
network_proxy: None,
|
|
}),
|
|
)
|
|
.await??;
|
|
|
|
let process = Arc::clone(&session.process);
|
|
let mut events = process.subscribe_events();
|
|
let mut output = Vec::new();
|
|
let mut last_seq = 0;
|
|
while !output.ends_with(b"\n") {
|
|
match timeout(Duration::from_secs(5), events.recv()).await?? {
|
|
ExecProcessEvent::Output(chunk) => {
|
|
assert_eq!(chunk.seq, last_seq + 1);
|
|
last_seq = chunk.seq;
|
|
output.extend_from_slice(&chunk.chunk.into_inner());
|
|
}
|
|
event => anyhow::bail!("expected ready output before disconnect, got {event:?}"),
|
|
}
|
|
}
|
|
let ready = String::from_utf8(output.clone())?;
|
|
let pid = ready
|
|
.strip_prefix("ready:")
|
|
.and_then(|line| line.strip_suffix('\n'))
|
|
.context("ready output should contain the process id")?
|
|
.to_string();
|
|
|
|
let mut connection_state = environment
|
|
.subscribe_connection_state()
|
|
.context("direct environment connection state")?;
|
|
assert_eq!(
|
|
*connection_state.borrow_and_update(),
|
|
EnvironmentConnectionState::Connected
|
|
);
|
|
proxy.pause_and_disconnect().await?;
|
|
timeout(
|
|
TEST_TIMEOUT,
|
|
connection_state.wait_for(|state| *state == EnvironmentConnectionState::Disconnected),
|
|
)
|
|
.await??;
|
|
tokio::fs::write(&gate_path, b"").await?;
|
|
timeout(Duration::from_secs(5), async {
|
|
while tokio::fs::metadata(&emitted_path).await.is_err() {
|
|
tokio::time::sleep(Duration::from_millis(10)).await;
|
|
}
|
|
})
|
|
.await
|
|
.context("process did not emit output while disconnected")?;
|
|
|
|
let process_for_read = Arc::clone(&process);
|
|
let mut pending_read = tokio::spawn(async move {
|
|
process_for_read
|
|
.read(
|
|
/*after_seq*/ Some(last_seq),
|
|
/*max_bytes*/ None,
|
|
/*wait_ms*/ Some(0),
|
|
)
|
|
.await
|
|
});
|
|
assert!(
|
|
timeout(Duration::from_millis(200), &mut pending_read)
|
|
.await
|
|
.is_err(),
|
|
"process reads should wait while recovery is in progress"
|
|
);
|
|
proxy.resume()?;
|
|
let replacement = timeout(TEST_TIMEOUT, accepted_sockets.recv())
|
|
.await?
|
|
.context("real Direct executor should reconnect")?;
|
|
timeout(
|
|
TEST_TIMEOUT,
|
|
manager.replace_accepted_websocket("environment-1", replacement),
|
|
)
|
|
.await??;
|
|
timeout(
|
|
TEST_TIMEOUT,
|
|
connection_state.wait_for(|state| *state == EnvironmentConnectionState::Connected),
|
|
)
|
|
.await??;
|
|
assert!(Arc::ptr_eq(
|
|
&environment,
|
|
&manager
|
|
.default_environment()
|
|
.context("same environment should remain installed")?,
|
|
));
|
|
assert_eq!(
|
|
timeout(TEST_TIMEOUT, environment.force_info()).await??,
|
|
EnvironmentInfo::local()
|
|
);
|
|
|
|
let recovered_read = timeout(Duration::from_secs(5), pending_read)
|
|
.await
|
|
.context("timed out waiting for a read after recovery")??;
|
|
let recovered_read = recovered_read?;
|
|
assert_eq!(recovered_read.failure, None);
|
|
let recovered_output = recovered_read
|
|
.chunks
|
|
.into_iter()
|
|
.flat_map(|chunk| chunk.chunk.into_inner())
|
|
.collect::<Vec<_>>();
|
|
assert_eq!(
|
|
String::from_utf8(recovered_output)?,
|
|
format!("during:{pid}\n")
|
|
);
|
|
|
|
let write = timeout(Duration::from_secs(5), process.write(b"hello\n".to_vec()))
|
|
.await
|
|
.context("timed out waiting for a write after recovery")??;
|
|
assert_eq!(write.status, WriteStatus::Accepted);
|
|
|
|
let mut saw_exit = false;
|
|
loop {
|
|
match timeout(Duration::from_secs(5), events.recv()).await?? {
|
|
ExecProcessEvent::Output(chunk) => {
|
|
assert_eq!(chunk.seq, last_seq + 1);
|
|
last_seq = chunk.seq;
|
|
output.extend_from_slice(&chunk.chunk.into_inner());
|
|
}
|
|
ExecProcessEvent::Exited { seq, exit_code, .. } => {
|
|
assert_eq!(seq, last_seq + 1);
|
|
assert_eq!(exit_code, 7);
|
|
last_seq = seq;
|
|
saw_exit = true;
|
|
}
|
|
ExecProcessEvent::Closed { seq } => {
|
|
assert!(saw_exit, "closed must be delivered after exit");
|
|
assert_eq!(seq, last_seq + 1);
|
|
break;
|
|
}
|
|
ExecProcessEvent::Failed(message) => {
|
|
anyhow::bail!("process recovery failed: {message}");
|
|
}
|
|
}
|
|
}
|
|
assert_eq!(
|
|
String::from_utf8(output)?,
|
|
format!("ready:{pid}\nduring:{pid}\nafter:{pid}:hello\n")
|
|
);
|
|
}
|
|
|
|
registry.verify().await;
|
|
let registrations = registry
|
|
.received_requests()
|
|
.await
|
|
.context("registration requests")?;
|
|
assert_eq!(
|
|
registrations[0].body,
|
|
br#"{"transport":"direct_jsonrpc_v1"}"#
|
|
);
|
|
|
|
let _ = shutdown_tx.send(());
|
|
timeout(TEST_TIMEOUT, executor_task).await???;
|
|
server_task.abort();
|
|
let _ = server_task.await;
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn accepted_websocket_environment_is_ready_immediately() -> Result<()> {
|
|
let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?;
|
|
let (mut socket, manager) =
|
|
connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?;
|
|
|
|
let status_task =
|
|
tokio::spawn(async move { manager.get_environment_status("environment-1").await });
|
|
let request = receive_jsonrpc(&mut socket).await?;
|
|
let JSONRPCMessage::Request(JSONRPCRequest { id, method, .. }) = request else {
|
|
anyhow::bail!("expected environment status request, got {request:?}");
|
|
};
|
|
assert_eq!(method, "environment/status");
|
|
send_jsonrpc(
|
|
&mut socket,
|
|
JSONRPCMessage::Response(JSONRPCResponse {
|
|
id,
|
|
result: serde_json::to_value(EnvironmentStatus {
|
|
status: EnvironmentStatusKind::Ready,
|
|
})?,
|
|
}),
|
|
)
|
|
.await?;
|
|
|
|
assert_eq!(
|
|
timeout(TEST_TIMEOUT, status_task).await??,
|
|
Some(EnvironmentObservedStatus::Ready)
|
|
);
|
|
|
|
server_task.abort();
|
|
let _ = server_task.await;
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn accepted_websocket_replacement_retires_old_socket_and_retries() -> Result<()> {
|
|
let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?;
|
|
let (mut first_socket, manager) =
|
|
connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?;
|
|
let (mut rejected_socket, rejected_websocket) =
|
|
connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?;
|
|
manager
|
|
.replace_accepted_websocket("environment-1", rejected_websocket)
|
|
.await?;
|
|
let previous_socket_event = timeout(TEST_TIMEOUT, first_socket.next())
|
|
.await
|
|
.context("the previous accepted websocket should be retired before replacement")?;
|
|
assert!(
|
|
matches!(
|
|
previous_socket_event,
|
|
None | Some(Ok(Message::Close(_))) | Some(Err(_))
|
|
),
|
|
"the previous accepted websocket should close before replacement: {previous_socket_event:?}"
|
|
);
|
|
let initialize = receive_jsonrpc(&mut rejected_socket).await?;
|
|
let JSONRPCMessage::Request(JSONRPCRequest { id, method, .. }) = initialize else {
|
|
anyhow::bail!("expected replacement initialize request, got {initialize:?}");
|
|
};
|
|
assert_eq!(method, "initialize");
|
|
|
|
let (_overlapping_socket, overlapping_websocket) =
|
|
connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?;
|
|
manager
|
|
.replace_accepted_websocket("environment-1", overlapping_websocket)
|
|
.await
|
|
.expect_err("an overlapping replacement should be rejected");
|
|
|
|
send_jsonrpc(
|
|
&mut rejected_socket,
|
|
JSONRPCMessage::Error(JSONRPCError {
|
|
id,
|
|
error: JSONRPCErrorError {
|
|
code: SESSION_ALREADY_ATTACHED_ERROR_CODE,
|
|
message: "session session-1 is already attached to another connection".to_string(),
|
|
data: None,
|
|
},
|
|
}),
|
|
)
|
|
.await?;
|
|
let rejected_socket_event = timeout(TEST_TIMEOUT, rejected_socket.next())
|
|
.await
|
|
.context("rejected replacement websocket should close")?;
|
|
assert!(
|
|
matches!(
|
|
rejected_socket_event,
|
|
None | Some(Ok(Message::Close(_))) | Some(Err(_))
|
|
),
|
|
"rejected replacement websocket should close: {rejected_socket_event:?}"
|
|
);
|
|
|
|
let (mut replacement_socket, replacement_websocket) =
|
|
connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?;
|
|
manager
|
|
.replace_accepted_websocket("environment-1", replacement_websocket)
|
|
.await?;
|
|
complete_initialize(&mut replacement_socket, "session-1", Some("session-1")).await?;
|
|
|
|
server_task.abort();
|
|
let _ = server_task.await;
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn accepted_websocket_reconnect_recovers_running_process_and_output() -> Result<()> {
|
|
let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?;
|
|
let (mut first_socket, manager) =
|
|
connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?;
|
|
let environment = manager
|
|
.default_environment()
|
|
.context("default environment should be installed")?;
|
|
let backend = environment.get_exec_backend();
|
|
let process_id = ProcessId::from("process-1");
|
|
let process_task = tokio::spawn({
|
|
let process_id = process_id.clone();
|
|
async move {
|
|
backend
|
|
.start(ExecParams {
|
|
metadata: Default::default(),
|
|
process_id,
|
|
argv: vec!["test-command".to_string()],
|
|
cwd: PathUri::parse("file:///workspace")?,
|
|
env_policy: None,
|
|
env: HashMap::new(),
|
|
tty: false,
|
|
pipe_stdin: false,
|
|
arg0: None,
|
|
sandbox: None,
|
|
enforce_managed_network: false,
|
|
managed_network: None,
|
|
network_proxy: None,
|
|
shell_snapshot: None,
|
|
})
|
|
.await
|
|
.map_err(anyhow::Error::from)
|
|
}
|
|
});
|
|
let request = receive_jsonrpc(&mut first_socket).await?;
|
|
let JSONRPCMessage::Request(JSONRPCRequest {
|
|
id, method, params, ..
|
|
}) = request
|
|
else {
|
|
anyhow::bail!("expected process start request, got {request:?}");
|
|
};
|
|
assert_eq!(method, "process/start");
|
|
assert_eq!(
|
|
serde_json::from_value::<ExecParams>(params.context("process params should exist")?)?
|
|
.process_id,
|
|
process_id
|
|
);
|
|
send_jsonrpc(
|
|
&mut first_socket,
|
|
JSONRPCMessage::Response(JSONRPCResponse {
|
|
id,
|
|
result: serde_json::to_value(ExecResponse {
|
|
process_id: process_id.clone(),
|
|
sandbox_type: None,
|
|
})?,
|
|
}),
|
|
)
|
|
.await?;
|
|
let process = timeout(TEST_TIMEOUT, process_task).await???.process;
|
|
|
|
first_socket.close(/*close_frame*/ None).await?;
|
|
let (mut replacement_socket, replacement_websocket) =
|
|
connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?;
|
|
manager
|
|
.replace_accepted_websocket("environment-1", replacement_websocket)
|
|
.await?;
|
|
complete_initialize(&mut replacement_socket, "session-1", Some("session-1")).await?;
|
|
|
|
let request = receive_jsonrpc(&mut replacement_socket).await?;
|
|
let JSONRPCMessage::Request(JSONRPCRequest {
|
|
id, method, params, ..
|
|
}) = request
|
|
else {
|
|
anyhow::bail!("expected recovery process read request, got {request:?}");
|
|
};
|
|
assert_eq!(method, "process/read");
|
|
assert_eq!(
|
|
serde_json::from_value::<ReadParams>(params.context("read params should exist")?)?,
|
|
ReadParams {
|
|
process_id: process_id.clone(),
|
|
after_seq: Some(0),
|
|
max_bytes: None,
|
|
wait_ms: Some(0),
|
|
}
|
|
);
|
|
send_jsonrpc(
|
|
&mut replacement_socket,
|
|
JSONRPCMessage::Response(JSONRPCResponse {
|
|
id,
|
|
result: serde_json::to_value(ReadResponse {
|
|
chunks: Vec::new(),
|
|
next_seq: 1,
|
|
exited: false,
|
|
exit_code: None,
|
|
closed: false,
|
|
failure: None,
|
|
sandbox_denied: false,
|
|
})?,
|
|
}),
|
|
)
|
|
.await?;
|
|
|
|
let read_task = tokio::spawn(async move {
|
|
process.read(Some(0), /*max_bytes*/ None, Some(0)).await
|
|
});
|
|
let request = receive_jsonrpc(&mut replacement_socket).await?;
|
|
let JSONRPCMessage::Request(JSONRPCRequest {
|
|
id, method, params, ..
|
|
}) = request
|
|
else {
|
|
anyhow::bail!("expected existing process read request, got {request:?}");
|
|
};
|
|
assert_eq!(method, "process/read");
|
|
assert_eq!(
|
|
serde_json::from_value::<ReadParams>(params.context("read params should exist")?)?,
|
|
ReadParams {
|
|
process_id,
|
|
after_seq: Some(0),
|
|
max_bytes: None,
|
|
wait_ms: Some(0),
|
|
}
|
|
);
|
|
let response = ReadResponse {
|
|
chunks: Vec::new(),
|
|
next_seq: 1,
|
|
exited: false,
|
|
exit_code: None,
|
|
closed: false,
|
|
failure: None,
|
|
sandbox_denied: false,
|
|
};
|
|
send_jsonrpc(
|
|
&mut replacement_socket,
|
|
JSONRPCMessage::Response(JSONRPCResponse {
|
|
id,
|
|
result: serde_json::to_value(&response)?,
|
|
}),
|
|
)
|
|
.await?;
|
|
assert_eq!(timeout(TEST_TIMEOUT, read_task).await???, response);
|
|
|
|
server_task.abort();
|
|
let _ = server_task.await;
|
|
Ok(())
|
|
}
|
|
|
|
async fn start_acceptor() -> Result<(
|
|
String,
|
|
mpsc::UnboundedReceiver<AcceptedSocket>,
|
|
JoinHandle<()>,
|
|
)> {
|
|
let listener = TcpListener::bind("127.0.0.1:0").await?;
|
|
let local_addr = listener.local_addr()?;
|
|
let (accepted_tx, accepted_rx) = mpsc::unbounded_channel();
|
|
let app = Router::new()
|
|
.route("/", any(accept_websocket))
|
|
.with_state(accepted_tx);
|
|
let server_task = tokio::spawn(async move {
|
|
let result = axum::serve(listener, app).await;
|
|
assert!(
|
|
result.is_ok(),
|
|
"accepted websocket test server should run: {result:?}"
|
|
);
|
|
});
|
|
Ok((format!("ws://{local_addr}/"), accepted_rx, server_task))
|
|
}
|
|
|
|
async fn accept_websocket(
|
|
websocket: WebSocketUpgrade,
|
|
State(accepted_tx): State<mpsc::UnboundedSender<AcceptedSocket>>,
|
|
) -> impl IntoResponse {
|
|
websocket.on_upgrade(move |websocket| async move {
|
|
let _ = accepted_tx.send(websocket);
|
|
})
|
|
}
|
|
|
|
async fn connect_executor(
|
|
websocket_url: &str,
|
|
accepted_sockets: &mut mpsc::UnboundedReceiver<AcceptedSocket>,
|
|
session_id: &str,
|
|
) -> Result<(
|
|
WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>,
|
|
EnvironmentManager,
|
|
)> {
|
|
let (mut websocket, _) = connect_async(websocket_url).await?;
|
|
let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv())
|
|
.await?
|
|
.context("accepted websocket channel should remain open")?;
|
|
let manager_task = tokio::spawn(EnvironmentManager::from_accepted_websocket(
|
|
"environment-1".to_string(),
|
|
accepted_websocket,
|
|
accepted_options(),
|
|
HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault),
|
|
));
|
|
complete_initialize(&mut websocket, session_id, /*resume_session_id*/ None).await?;
|
|
let manager = timeout(TEST_TIMEOUT, manager_task).await???;
|
|
Ok((websocket, manager))
|
|
}
|
|
|
|
async fn connect_replacement_executor(
|
|
websocket_url: &str,
|
|
accepted_sockets: &mut mpsc::UnboundedReceiver<AcceptedSocket>,
|
|
) -> Result<(
|
|
WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>,
|
|
AcceptedSocket,
|
|
)> {
|
|
let (websocket, _) = connect_async(websocket_url).await?;
|
|
let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv())
|
|
.await?
|
|
.context("accepted websocket channel should remain open")?;
|
|
Ok((websocket, accepted_websocket))
|
|
}
|
|
|
|
fn accepted_options() -> ExecServerClientConnectOptions {
|
|
ExecServerClientConnectOptions {
|
|
client_name: "host-test".to_string(),
|
|
initialize_timeout: TEST_TIMEOUT,
|
|
resume_session_id: None,
|
|
}
|
|
}
|
|
|
|
async fn complete_initialize<S>(
|
|
websocket: &mut WebSocketStream<S>,
|
|
session_id: &str,
|
|
resume_session_id: Option<&str>,
|
|
) -> Result<()>
|
|
where
|
|
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
|
|
{
|
|
let initialize = receive_jsonrpc(&mut *websocket).await?;
|
|
let JSONRPCMessage::Request(JSONRPCRequest {
|
|
id, method, params, ..
|
|
}) = initialize
|
|
else {
|
|
anyhow::bail!("expected initialize request, got {initialize:?}");
|
|
};
|
|
assert_eq!(method, "initialize");
|
|
assert_eq!(
|
|
serde_json::from_value::<InitializeParams>(
|
|
params.context("initialize request should contain params")?
|
|
)?,
|
|
InitializeParams {
|
|
client_name: "host-test".to_string(),
|
|
resume_session_id: resume_session_id.map(str::to_string),
|
|
}
|
|
);
|
|
send_jsonrpc(
|
|
&mut *websocket,
|
|
JSONRPCMessage::Response(JSONRPCResponse {
|
|
id,
|
|
result: serde_json::to_value(InitializeResponse {
|
|
session_id: session_id.to_string(),
|
|
environment_info: Some(EnvironmentInfo::local()),
|
|
})?,
|
|
}),
|
|
)
|
|
.await?;
|
|
let initialized = receive_jsonrpc(&mut *websocket).await?;
|
|
assert_eq!(
|
|
initialized,
|
|
JSONRPCMessage::Notification(JSONRPCNotification {
|
|
method: "initialized".to_string(),
|
|
params: Some(serde_json::json!({})),
|
|
})
|
|
);
|
|
Ok(())
|
|
}
|
|
|
|
async fn send_jsonrpc<S>(websocket: &mut WebSocketStream<S>, message: JSONRPCMessage) -> Result<()>
|
|
where
|
|
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
|
|
{
|
|
websocket
|
|
.send(Message::Text(serde_json::to_string(&message)?.into()))
|
|
.await?;
|
|
Ok(())
|
|
}
|
|
|
|
async fn receive_jsonrpc<S>(websocket: &mut WebSocketStream<S>) -> Result<JSONRPCMessage>
|
|
where
|
|
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
|
|
{
|
|
loop {
|
|
let message = websocket
|
|
.next()
|
|
.await
|
|
.context("accepted websocket should remain open")??;
|
|
if let Message::Text(text) = message {
|
|
return Ok(serde_json::from_str(&text)?);
|
|
}
|
|
}
|
|
}
|