use super::subprocess::{ETXTBSY_RETRY_ATTEMPTS, SpawnArgs, spawn_and_capture};
use super::{DEFAULT_TOOL_DEADLINE, ENV_CONV_BRANCH, ENV_CONV_REPO, ExecError};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::ffi::OsString;
use std::os::unix::process::ExitStatusExt;
use std::path::Path;
use std::sync::atomic::AtomicBool;
use thiserror::Error;
#[derive(Debug, Serialize)]
pub struct ControlRequest<'a> {
pub id: &'a str,
pub name: &'a str,
pub input: &'a Value,
pub role: &'a str,
pub agent_id: &'a str,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Verdict {
Pass,
Refuse { reason: String },
Hold { reason: String },
}
#[derive(Debug, Error)]
pub enum ControlError {
#[error("spawn control {command:?}: {source}")]
Spawn {
command: String,
#[source]
source: std::io::Error,
},
#[error("control {command:?} terminated by signal {signal}")]
KilledBySignal { command: String, signal: i32 },
#[error("control {command:?} broke the verdict protocol: {detail}")]
Protocol { command: String, detail: String },
}
pub fn consult(
command: &str,
request: &ControlRequest<'_>,
conv_repo: &Path,
stop: &AtomicBool,
) -> Result<Verdict, ControlError> {
let stdin = serde_json::to_vec(request).expect("ControlRequest serializes");
let binary = OsString::from(command);
let extra_env = [
(ENV_CONV_REPO, conv_repo.as_os_str().to_owned()),
(ENV_CONV_BRANCH, OsString::from(request.agent_id)),
];
let captured = spawn_and_capture(&SpawnArgs {
binary: &binary,
args: &[],
stdin_bytes: &stdin,
extra_env: &extra_env,
cwd: conv_repo,
stop,
deadline: DEFAULT_TOOL_DEADLINE,
etxtbsy_budget: ETXTBSY_RETRY_ATTEMPTS,
tool_name: command,
})
.map_err(|e| spawn_fault(command, e))?;
let Some(code) = captured.status.code() else {
return Err(ControlError::KilledBySignal {
command: command.to_string(),
signal: captured.status.signal().unwrap_or(0),
});
};
if code != 0 {
return Err(ControlError::Protocol {
command: command.to_string(),
detail: format!(
"exited {code}; stderr: {}",
String::from_utf8_lossy(&captured.stderr).trim()
),
});
}
parse_verdict(&captured.stdout).map_err(|detail| ControlError::Protocol {
command: command.to_string(),
detail: format!(
"stdout is not a verdict ({detail}): {}",
String::from_utf8_lossy(&captured.stdout).trim()
),
})
}
fn spawn_fault(command: &str, e: ExecError) -> ControlError {
match e {
ExecError::Spawn { source, .. } => ControlError::Spawn {
command: command.to_string(),
source,
},
other => ControlError::Protocol {
command: command.to_string(),
detail: other.to_string(),
},
}
}
fn parse_verdict(stdout: &[u8]) -> Result<Verdict, String> {
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct Raw {
verdict: String,
#[serde(default)]
reason: Option<String>,
}
let raw: Raw = serde_json::from_slice(stdout).map_err(|e| e.to_string())?;
match (raw.verdict.as_str(), raw.reason) {
("pass", None) => Ok(Verdict::Pass),
("pass", Some(_)) => Err("a pass carries no reason".into()),
("refuse", Some(reason)) => Ok(Verdict::Refuse { reason }),
("hold", Some(reason)) => Ok(Verdict::Hold { reason }),
(v @ ("refuse" | "hold"), None) => Err(format!("{v:?} requires a reason")),
(other, _) => Err(format!("unknown verdict {other:?}")),
}
}
#[cfg(test)]
mod tests;