use std::{path::PathBuf, process::Stdio, time::Duration};
use serde::{Deserialize, Serialize};
use tokio::io::AsyncReadExt as _;
use crate::configuration::tokens::ExternallyManaged;
use super::secret_string::SecretAccessToken;
const MAX_PIPE_BYTES: usize = 4 * 1024;
const DEFAULT_TIMEOUT_SECONDS: u64 = 30;
const fn default_timeout_seconds() -> u64 {
DEFAULT_TIMEOUT_SECONDS
}
#[allow(clippy::trivially_copy_pass_by_ref, reason = "serde needs a reference")]
fn is_default_timeout_seconds(timeout_seconds: &u64) -> bool {
timeout_seconds == &DEFAULT_TIMEOUT_SECONDS
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct ExternallyManagedCredential {
pub command: PathBuf,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub args: Vec<String>,
#[serde(
default = "default_timeout_seconds",
skip_serializing_if = "is_default_timeout_seconds"
)]
pub timeout_seconds: u64,
}
impl From<ExternallyManagedCredential> for ExternallyManaged {
fn from(credential: ExternallyManagedCredential) -> Self {
Self::from_async(move |_auth_server| {
let credential = credential.clone();
async move {
credential
.request_access_token()
.await
.map(|token| token.secret().to_string())
.map_err(Into::into)
}
})
}
}
impl ExternallyManagedCredential {
pub async fn request_access_token(&self) -> Result<SecretAccessToken, ExternalCommandError> {
let output = Box::pin(self.run()).await?;
let token = String::from_utf8(output)
.map_err(|_| ExternalCommandError::InvalidUtf8 {
program: self.command.clone(),
})?
.trim()
.to_string();
if token.is_empty() {
return Err(ExternalCommandError::EmptyOutput {
program: self.command.clone(),
});
}
Ok(SecretAccessToken::from(token))
}
async fn run(&self) -> Result<Vec<u8>, ExternalCommandError> {
let mut command = tokio::process::Command::new(&self.command);
command
.args(&self.args)
.stdin(Stdio::null())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.kill_on_drop(true);
let mut child = command
.spawn()
.map_err(|source| ExternalCommandError::Spawn {
program: self.command.clone(),
source,
})?;
let stdout = child.stdout.take().expect("stdout is piped");
let stderr = child.stderr.take().expect("stderr is piped");
let timeout = Duration::from_secs(self.timeout_seconds);
let (status, stdout_buf, stderr_buf, stdout_truncated) =
tokio::time::timeout(timeout, async {
let (stdout_result, stderr_result) = futures::future::join(
read_capped(stdout, MAX_PIPE_BYTES),
read_capped(stderr, MAX_PIPE_BYTES),
)
.await;
let (stdout_buf, stdout_truncated) =
stdout_result.map_err(|source| ExternalCommandError::Read {
program: self.command.clone(),
source,
})?;
let (stderr_buf, _) = stderr_result.unwrap_or_default();
let status = child
.wait()
.await
.map_err(|source| ExternalCommandError::Read {
program: self.command.clone(),
source,
})?;
Ok::<_, ExternalCommandError>((status, stdout_buf, stderr_buf, stdout_truncated))
})
.await
.map_err(|_| ExternalCommandError::Timeout {
program: self.command.clone(),
timeout,
})??;
if !status.success() {
return Err(ExternalCommandError::ExitStatus {
program: self.command.clone(),
status: status.to_string(),
stderr: String::from_utf8_lossy(&stderr_buf).trim().to_string(),
});
}
if stdout_truncated {
return Err(ExternalCommandError::OutputTooLarge {
program: self.command.clone(),
limit: MAX_PIPE_BYTES,
});
}
Ok(stdout_buf)
}
}
async fn read_capped(
mut reader: impl tokio::io::AsyncRead + Unpin,
cap: usize,
) -> std::io::Result<(Vec<u8>, bool)> {
let mut retained = Vec::new();
let mut chunk = vec![0_u8; 1024];
let mut truncated = false;
loop {
let read = reader.read(&mut chunk).await?;
if read == 0 {
return Ok((retained, truncated));
}
let room = cap - retained.len();
if read > room {
truncated = true;
}
retained.extend_from_slice(&chunk[..read.min(room)]);
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum ExternalCommandError {
#[error("failed to run {program:?}: {source}")]
Spawn {
program: PathBuf,
source: std::io::Error,
},
#[error("failed to read the output of {program:?}: {source}")]
Read {
program: PathBuf,
source: std::io::Error,
},
#[error("{program:?} did not produce an access token within {timeout:?}")]
Timeout {
program: PathBuf,
timeout: Duration,
},
#[error("{program:?} failed with {status}: {stderr}")]
ExitStatus {
program: PathBuf,
status: String,
stderr: String,
},
#[error("{program:?} wrote more than {limit} bytes to stdout")]
OutputTooLarge {
program: PathBuf,
limit: usize,
},
#[error("{program:?} did not write a valid UTF-8 access token to stdout")]
InvalidUtf8 {
program: PathBuf,
},
#[error("{program:?} did not write an access token to stdout")]
EmptyOutput {
program: PathBuf,
},
}
#[cfg(test)]
pub(super) use tests::shell;
#[cfg(test)]
mod tests {
use std::path::PathBuf;
use super::{DEFAULT_TIMEOUT_SECONDS, ExternalCommandError, ExternallyManagedCredential};
fn credential(command: impl Into<PathBuf>) -> ExternallyManagedCredential {
ExternallyManagedCredential {
command: command.into(),
args: Vec::new(),
timeout_seconds: DEFAULT_TIMEOUT_SECONDS,
}
}
pub(in super::super) fn shell() -> (PathBuf, &'static str) {
#[cfg(windows)]
{
let comspec = std::env::var_os("COMSPEC")
.unwrap_or_else(|| r"C:\Windows\System32\cmd.exe".into());
(PathBuf::from(comspec), "/C")
}
#[cfg(not(windows))]
{
(PathBuf::from("/bin/sh"), "-c")
}
}
fn script_credential(script: &str) -> ExternallyManagedCredential {
let (program, flag) = shell();
ExternallyManagedCredential {
args: vec![flag.to_string(), script.to_string()],
..credential(program)
}
}
fn echo_credential(token: &str) -> ExternallyManagedCredential {
#[cfg(windows)]
let script = format!("echo {token}");
#[cfg(not(windows))]
let script = format!("printf '%s\\n' '{token}'");
script_credential(&script)
}
#[tokio::test]
async fn returns_trimmed_stdout_as_the_access_token() {
let token = echo_credential("an-access-token")
.request_access_token()
.await
.expect("the command should produce a token");
assert_eq!(token.secret(), "an-access-token");
}
#[tokio::test]
async fn reports_stderr_when_the_command_fails() {
#[cfg(windows)]
let script = "echo something went wrong 1>&2 && exit 3";
#[cfg(not(windows))]
let script = "echo 'something went wrong' >&2; exit 3";
let error = script_credential(script)
.request_access_token()
.await
.expect_err("a failing command should be an error");
let message = error.to_string();
assert!(
message.contains("something went wrong"),
"stderr should be reported: {message}"
);
}
#[tokio::test]
async fn times_out_a_command_that_hangs() {
#[cfg(windows)]
let script = "ping -n 30 127.0.0.1 > nul";
#[cfg(not(windows))]
let script = "sleep 30";
let error = ExternallyManagedCredential {
timeout_seconds: 1,
..script_credential(script)
}
.request_access_token()
.await
.expect_err("a hanging command should time out");
assert!(
matches!(error, ExternalCommandError::Timeout { .. }),
"unexpected error: {error}"
);
}
#[tokio::test]
async fn rejects_empty_output() {
let error = script_credential("exit 0")
.request_access_token()
.await
.expect_err("a command that prints nothing should be an error");
assert!(
matches!(error, ExternalCommandError::EmptyOutput { .. }),
"unexpected error: {error}"
);
}
#[tokio::test]
async fn rejects_output_that_is_too_large() {
#[cfg(windows)]
let script = "for /L %i in (1,1,20000) do @echo aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa";
#[cfg(not(windows))]
let script = "yes aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa | head -c 200000";
let error = script_credential(script)
.request_access_token()
.await
.expect_err("an oversized output should be an error");
assert!(
matches!(error, ExternalCommandError::OutputTooLarge { .. }),
"unexpected error: {error}"
);
}
}