use std::collections::BTreeMap;
use std::time::Duration;
use async_trait::async_trait;
use base64::engine::general_purpose::STANDARD as BASE64;
use base64::Engine as _;
use futures::stream::BoxStream;
use futures::StreamExt;
use serde::Deserialize;
use serde_json::json;
use crate::error::{ErrorData, Result};
use crate::traits::{CommandOutput, RunCommandRequest};
use alien_error::{AlienError, Context, IntoAlienError};
pub use alien_core::sandbox_process::AGENT_PORT;
const RUN_COMMAND: &str = "sandbox.runCommand";
#[cfg(not(test))]
const AGENT_RESPONSE_TIMEOUT: Duration = Duration::from_secs(60);
#[cfg(test)]
const AGENT_RESPONSE_TIMEOUT: Duration = Duration::from_millis(200);
#[async_trait]
pub trait AgentTransport: Send + Sync + std::fmt::Debug {
async fn request(
&self,
session_id: &str,
method: reqwest::Method,
path: &str,
) -> Result<reqwest::RequestBuilder>;
fn provider(&self) -> &'static str;
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase", tag = "t")]
enum AgentFrame {
Stdout {
seq: u64,
data: String,
},
Stderr {
seq: u64,
data: String,
},
Exit {
code: i32,
#[serde(default)]
truncated: bool,
},
Error {
code: String,
message: String,
},
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct ReadFileResponse {
contents_base64: String,
}
pub async fn run_command<T: AgentTransport + ?Sized>(
transport: &T,
session_id: &str,
request: RunCommandRequest,
) -> Result<BoxStream<'static, Result<CommandOutput>>> {
if deadline_millis(request.deadline) == 0 {
return Err(AlienError::new(ErrorData::SandboxCommandFailed {
failure: "invalidRequest".to_string(),
reason: "a command must carry a non-zero deadline".to_string(),
}));
}
let body = json!({
"command": request.command,
"deadlineMs": deadline_millis(request.deadline),
"workingDirectory": request.working_directory,
"env": request.env,
});
let response = send(
transport
.request(session_id, reqwest::Method::POST, "/v1/exec")
.await?
.json(&body),
RUN_COMMAND,
)
.await?;
Ok(frame_stream(response, transport.provider()))
}
pub async fn read_file<T: AgentTransport + ?Sized>(
transport: &T,
session_id: &str,
path: &str,
) -> Result<Vec<u8>> {
let response = send(
transport
.request(session_id, reqwest::Method::GET, "/v1/files")
.await?
.query(&[("path", path)]),
"sandbox.readFile",
)
.await?;
let body: ReadFileResponse =
response
.json()
.await
.into_alien_error()
.context(ErrorData::UnexpectedResponseFormat {
provider: transport.provider().to_string(),
binding_name: "sandbox.readFile".to_string(),
field: "body".to_string(),
response_json: "the agent returned a body this provider cannot parse".to_string(),
})?;
decode(
&body.contents_base64,
transport.provider(),
"sandbox.readFile",
"contentsBase64",
)
}
pub async fn write_files<T: AgentTransport + ?Sized>(
transport: &T,
session_id: &str,
files: BTreeMap<String, Vec<u8>>,
) -> Result<()> {
for (path, contents) in files {
send(
transport
.request(session_id, reqwest::Method::PUT, "/v1/files")
.await?
.json(&json!({
"path": path,
"contentsBase64": BASE64.encode(contents),
})),
"sandbox.writeFiles",
)
.await?;
}
Ok(())
}
pub async fn mkdir<T: AgentTransport + ?Sized>(
transport: &T,
session_id: &str,
path: &str,
) -> Result<()> {
send(
transport
.request(session_id, reqwest::Method::POST, "/v1/mkdir")
.await?
.json(&json!({ "path": path })),
"sandbox.mkdir",
)
.await?;
Ok(())
}
fn deadline_millis(deadline: Duration) -> u64 {
u64::try_from(deadline.as_millis()).unwrap_or(u64::MAX)
}
fn unanswered(operation: &str, reason: &str) -> ErrorData {
if operation == RUN_COMMAND {
return ErrorData::SandboxCommandFailed {
failure: "outcomeUnknown".to_string(),
reason: format!("{reason}; the command may have started, its outcome is unknown"),
};
}
ErrorData::SandboxUnreachable {
operation: operation.to_string(),
reason: reason.to_string(),
}
}
pub async fn send(request: reqwest::RequestBuilder, operation: &str) -> Result<reqwest::Response> {
let response = match tokio::time::timeout(AGENT_RESPONSE_TIMEOUT, request.send()).await {
Ok(sent) => sent
.into_alien_error()
.context(unanswered(operation, "the request never reached the agent"))?,
Err(_) => {
return Err(AlienError::new(unanswered(
operation,
&format!(
"the agent did not answer within {}s",
AGENT_RESPONSE_TIMEOUT.as_secs()
),
)));
}
};
if response.status().is_success() {
return Ok(response);
}
let status = response.status();
let body = match response.text().await {
Ok(body) => body,
Err(error) if operation == RUN_COMMAND => {
return Err(error)
.into_alien_error()
.context(ErrorData::SandboxCommandFailed {
failure: "outcomeUnknown".to_string(),
reason: format!(
"{operation} returned {status} and its body could not be read, so whether \
it ran is unknown"
),
})
}
Err(error) => {
return Err(error)
.into_alien_error()
.context(ErrorData::SandboxUnreachable {
operation: operation.to_string(),
reason: format!("{operation} returned {status} and its body could not be read"),
})
}
};
if status.is_server_error() && body.trim().is_empty() {
if operation == RUN_COMMAND {
return Err(AlienError::new(ErrorData::SandboxCommandFailed {
failure: "outcomeUnknown".to_string(),
reason: format!(
"the sandbox host returned {status} with no response from the agent, so \
whether the command ran is unknown"
),
}));
}
return Err(AlienError::new(ErrorData::SandboxUnreachable {
operation: operation.to_string(),
reason: format!(
"the sandbox host returned {status} before the request reached the agent"
),
}));
}
if status.is_server_error() && operation != RUN_COMMAND {
return Err(AlienError::new(ErrorData::SandboxUnreachable {
operation: operation.to_string(),
reason: format!("the agent returned {status}: {body}"),
}));
}
Err(AlienError::new(ErrorData::SandboxCommandFailed {
failure: "agentRefused".to_string(),
reason: format!("{operation} returned {status}: {body}"),
}))
}
fn frame_stream(
response: reqwest::Response,
provider: &'static str,
) -> BoxStream<'static, Result<CommandOutput>> {
struct State {
bytes: BoxStream<'static, reqwest::Result<bytes::Bytes>>,
buffer: Vec<u8>,
finished: bool,
saw_terminal: bool,
provider: &'static str,
}
let state = State {
bytes: response.bytes_stream().boxed(),
buffer: Vec::new(),
finished: false,
saw_terminal: false,
provider,
};
futures::stream::unfold(state, |mut state| async move {
loop {
if let Some(index) = state.buffer.iter().position(|byte| *byte == b'\n') {
let line: Vec<u8> = state.buffer.drain(..=index).collect();
let line = &line[..line.len() - 1];
if line.is_empty() {
continue;
}
let frame = match serde_json::from_slice::<AgentFrame>(line) {
Ok(frame) => frame,
Err(error) => {
state.finished = true;
let failure = malformed(&error.to_string(), state.provider);
return Some((Err(failure), state));
}
};
if matches!(frame, AgentFrame::Exit { .. } | AgentFrame::Error { .. }) {
state.saw_terminal = true;
}
let output = frame.into_output(state.provider);
return Some((output, state));
}
if state.finished {
return None;
}
match state.bytes.next().await {
Some(Ok(chunk)) => state.buffer.extend_from_slice(&chunk),
Some(Err(error)) => {
state.finished = true;
return Some((
Err(AlienError::new(unanswered(
RUN_COMMAND,
&format!("the output stream failed: {error}"),
))),
state,
));
}
None => {
state.finished = true;
if !state.saw_terminal {
return Some((
Err(AlienError::new(unanswered(
RUN_COMMAND,
"the output stream ended without a terminal frame",
))),
state,
));
}
return None;
}
}
}
})
.boxed()
}
impl AgentFrame {
fn into_output(self, provider: &'static str) -> Result<CommandOutput> {
match self {
Self::Stdout { seq, data } => Ok(CommandOutput::Stdout {
seq,
data: decode(&data, provider, RUN_COMMAND, "data")?,
}),
Self::Stderr { seq, data } => Ok(CommandOutput::Stderr {
seq,
data: decode(&data, provider, RUN_COMMAND, "data")?,
}),
Self::Exit { code, truncated } => Ok(CommandOutput::Exit { code, truncated }),
Self::Error { code, message } => {
Err(AlienError::new(ErrorData::SandboxCommandFailed {
failure: code,
reason: message,
}))
}
}
}
}
fn decode(data: &str, provider: &'static str, binding_name: &str, field: &str) -> Result<Vec<u8>> {
BASE64
.decode(data)
.into_alien_error()
.context(ErrorData::UnexpectedResponseFormat {
provider: provider.to_string(),
binding_name: binding_name.to_string(),
field: field.to_string(),
response_json: format!("{field} was not valid base64"),
})
}
fn malformed(reason: &str, provider: &'static str) -> AlienError<ErrorData> {
AlienError::new(ErrorData::UnexpectedResponseFormat {
provider: provider.to_string(),
binding_name: RUN_COMMAND.to_string(),
field: "frame".to_string(),
response_json: format!("an output frame did not parse: {reason}"),
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::traits::CommandOutput;
use axum::http::StatusCode;
use axum::response::IntoResponse;
use axum::routing::post;
use axum::Router;
use std::net::SocketAddr;
async fn serve_frames(chunks: Vec<&'static str>) -> String {
let handler = move || {
let chunks = chunks.clone();
async move {
let stream = futures::stream::iter(
chunks
.into_iter()
.map(|chunk| Ok::<_, std::io::Error>(bytes::Bytes::from(chunk))),
);
axum::body::Body::from_stream(stream).into_response()
}
};
let router = Router::new().route("/v1/exec", post(handler));
let listener = tokio::net::TcpListener::bind::<SocketAddr>("127.0.0.1:0".parse().unwrap())
.await
.expect("bind");
let address = listener.local_addr().expect("address");
tokio::spawn(async move {
axum::serve(listener, router).await.expect("serve");
});
format!("http://{address}")
}
async fn frames_from(chunks: Vec<&'static str>) -> Vec<Result<CommandOutput>> {
let base = serve_frames(chunks).await;
let response = reqwest::Client::new()
.post(format!("{base}/v1/exec"))
.send()
.await
.expect("responds");
frame_stream(response, "test-sandbox")
.collect::<Vec<_>>()
.await
}
async fn send_status(status: StatusCode, body: &'static str) -> AlienError<ErrorData> {
let handler = move || async move { (status, body).into_response() };
let router = Router::new().route("/v1/exec", post(handler));
let listener = tokio::net::TcpListener::bind::<SocketAddr>("127.0.0.1:0".parse().unwrap())
.await
.expect("bind");
let address = listener.local_addr().expect("address");
tokio::spawn(async move { axum::serve(listener, router).await.expect("serve") });
send(
reqwest::Client::new().post(format!("http://{address}/v1/exec")),
RUN_COMMAND,
)
.await
.expect_err("a non-success must be an error")
}
#[tokio::test]
async fn a_bodyless_server_error_is_not_reported_as_the_agent_refusing() {
let error = send_status(StatusCode::BAD_GATEWAY, "").await;
let rendered = error.to_string();
assert!(
!rendered.contains("agentRefused"),
"a proxy 502 is not the agent refusing: {rendered}"
);
assert!(
rendered.contains("outcomeUnknown"),
"a proxy can synthesize a 502 after the agent accepted the request, so the caller has \
to be told the outcome is unknown rather than that it is safe to repeat: {rendered}"
);
}
#[tokio::test]
async fn a_server_error_carrying_the_agents_own_cause_stays_a_refusal() {
let error = send_status(StatusCode::INTERNAL_SERVER_ERROR, "spawn failed: ENOMEM").await;
let rendered = error.to_string();
assert!(
rendered.contains("spawn failed: ENOMEM"),
"the agent's cause has to survive: {rendered}"
);
assert!(
!rendered.contains("outcomeUnknown"),
"the agent answered with its own cause, so the outcome is not unknown: {rendered}"
);
}
#[tokio::test]
async fn an_agent_that_never_answers_is_refused_within_the_bound() {
let handler = || async {
tokio::time::sleep(AGENT_RESPONSE_TIMEOUT * 20).await;
"late".into_response()
};
let router = Router::new().route("/v1/exec", post(handler));
let listener = tokio::net::TcpListener::bind::<SocketAddr>("127.0.0.1:0".parse().unwrap())
.await
.expect("bind");
let address = listener.local_addr().expect("address");
tokio::spawn(async move { axum::serve(listener, router).await.expect("serve") });
let started = std::time::Instant::now();
let error = send(
reqwest::Client::new().post(format!("http://{address}/v1/exec")),
RUN_COMMAND,
)
.await
.expect_err("a stalled agent must be refused, not waited on");
assert!(
started.elapsed() < AGENT_RESPONSE_TIMEOUT * 10,
"refused at the bound, not at the agent's leisure: {:?}",
started.elapsed()
);
assert!(
error.to_string().contains("did not answer"),
"the refusal says the agent stalled: {error}"
);
}
#[tokio::test]
async fn a_slow_body_after_prompt_headers_is_not_cut_off() {
let handler = || async {
let frames = async_stream_frames(vec![
"{\"t\":\"stdout\",\"seq\":0,\"data\":\"aGk=\"}\n",
"{\"t\":\"exit\",\"code\":0,\"truncated\":false}\n",
]);
axum::body::Body::from_stream(frames).into_response()
};
let router = Router::new().route("/v1/exec", post(handler));
let listener = tokio::net::TcpListener::bind::<SocketAddr>("127.0.0.1:0".parse().unwrap())
.await
.expect("bind");
let address = listener.local_addr().expect("address");
tokio::spawn(async move { axum::serve(listener, router).await.expect("serve") });
let response = send(
reqwest::Client::new().post(format!("http://{address}/v1/exec")),
RUN_COMMAND,
)
.await
.expect("headers arrive at once");
let outputs = frame_stream(response, "test-sandbox")
.collect::<Vec<_>>()
.await;
assert_eq!(outputs.len(), 2, "every frame arrived: {outputs:?}");
assert_eq!(
outputs[1].as_ref().expect("exit"),
&CommandOutput::Exit {
code: 0,
truncated: false
}
);
}
fn async_stream_frames(
chunks: Vec<&'static str>,
) -> impl futures::Stream<Item = std::result::Result<&'static str, std::io::Error>> {
futures::stream::iter(chunks).then(|chunk| async move {
tokio::time::sleep(AGENT_RESPONSE_TIMEOUT * 2).await;
Ok(chunk)
})
}
async fn stalled_agent() -> String {
let handler = || async {
tokio::time::sleep(AGENT_RESPONSE_TIMEOUT * 20).await;
"late".into_response()
};
let router = Router::new().fallback(handler);
let listener = tokio::net::TcpListener::bind::<SocketAddr>("127.0.0.1:0".parse().unwrap())
.await
.expect("bind");
let address = listener.local_addr().expect("address");
tokio::spawn(async move { axum::serve(listener, router).await.expect("serve") });
format!("http://{address}")
}
#[tokio::test]
async fn a_dropped_run_command_connection_is_not_retryable() {
let listener = tokio::net::TcpListener::bind::<SocketAddr>("127.0.0.1:0".parse().unwrap())
.await
.expect("bind");
let address = listener.local_addr().expect("address");
tokio::spawn(async move {
loop {
if let Ok((socket, _)) = listener.accept().await {
drop(socket);
}
}
});
let error = send(
reqwest::Client::new().post(format!("http://{address}/v1/exec")),
RUN_COMMAND,
)
.await
.expect_err("a dropped connection is a refusal, not a response");
assert_eq!(error.code, "SANDBOX_COMMAND_FAILED", "got: {error}");
assert!(
!error.retryable,
"a command that may have started must not be retried: {error}"
);
assert!(
error.to_string().contains("may have started"),
"the refusal says the outcome is unknown: {error}"
);
}
#[tokio::test]
async fn a_stalled_run_command_is_not_retryable() {
let base = stalled_agent().await;
let error = send(
reqwest::Client::new().post(format!("{base}/v1/exec")),
RUN_COMMAND,
)
.await
.expect_err("a stalled agent must be refused, not waited on");
assert_eq!(error.code, "SANDBOX_COMMAND_FAILED", "got: {error}");
assert!(
!error.retryable,
"a command with an unknown outcome must not be retried: {error}"
);
assert!(
error.to_string().contains("did not answer")
&& error.to_string().contains("may have started"),
"the refusal says the agent stalled and the outcome is unknown: {error}"
);
}
#[tokio::test]
async fn a_stalled_file_operation_stays_retryable() {
let base = stalled_agent().await;
let error = send(
reqwest::Client::new().get(format!("{base}/v1/files")),
"sandbox.readFile",
)
.await
.expect_err("a stalled agent must be refused, not waited on");
assert_eq!(error.code, "SANDBOX_UNREACHABLE", "got: {error}");
assert!(
error.retryable,
"a stalled file read is safe to repeat: {error}"
);
assert!(
error.to_string().contains("did not answer"),
"the refusal says the agent stalled: {error}"
);
}
#[tokio::test]
async fn frames_decode_in_order_with_a_real_exit_code() {
let outputs = frames_from(vec![
"{\"t\":\"stdout\",\"seq\":0,\"data\":\"aGk=\"}\n",
"{\"t\":\"stderr\",\"seq\":1,\"data\":\"b29wcw==\"}\n",
"{\"t\":\"exit\",\"code\":7,\"truncated\":false}\n",
])
.await;
assert_eq!(outputs.len(), 3);
assert_eq!(
outputs[0].as_ref().expect("stdout"),
&CommandOutput::Stdout {
seq: 0,
data: b"hi".to_vec()
}
);
assert_eq!(
outputs[1].as_ref().expect("stderr"),
&CommandOutput::Stderr {
seq: 1,
data: b"oops".to_vec()
}
);
assert_eq!(
outputs[2].as_ref().expect("exit"),
&CommandOutput::Exit {
code: 7,
truncated: false
}
);
}
#[tokio::test]
async fn a_frame_split_across_chunks_is_reassembled() {
let outputs = frames_from(vec![
"{\"t\":\"stdo",
"ut\",\"seq\":0,\"data\":\"aGk=\"}\n{\"t\":\"ex",
"it\",\"code\":0,\"truncated\":false}\n",
])
.await;
assert_eq!(
outputs.len(),
2,
"a split frame must not become two frames or an error"
);
assert_eq!(
outputs[0].as_ref().expect("stdout"),
&CommandOutput::Stdout {
seq: 0,
data: b"hi".to_vec()
}
);
assert_eq!(
outputs[1].as_ref().expect("exit"),
&CommandOutput::Exit {
code: 0,
truncated: false
}
);
}
#[tokio::test]
async fn a_stream_without_a_terminal_frame_is_an_unknown_outcome() {
let outputs = frames_from(vec!["{\"t\":\"stdout\",\"seq\":0,\"data\":\"aGk=\"}\n"]).await;
assert_eq!(outputs.len(), 2);
outputs[0].as_ref().expect("the stdout frame still arrives");
let error = outputs[1]
.as_ref()
.expect_err("a truncated stream must not read as success");
assert!(
error.to_string().contains("without a terminal frame"),
"the failure must name the cause: {error}"
);
assert_eq!(error.code, "SANDBOX_COMMAND_FAILED", "got: {error}");
assert!(
!error.retryable,
"a command that started and whose end was lost must not be retried: {error}"
);
}
#[tokio::test]
async fn an_error_frame_surfaces_as_an_error_not_a_silent_end() {
let outputs = frames_from(vec![
"{\"t\":\"error\",\"code\":\"deadlineExceeded\",\"message\":\"exceeded its 300ms deadline\"}\n",
])
.await;
assert_eq!(outputs.len(), 1);
let error = outputs[0]
.as_ref()
.expect_err("an error frame is a failure");
assert!(error.to_string().contains("deadlineExceeded"), "{error}");
}
}