Skip to main content

azure_identity/process/
mod.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4// cspell:ignore workdir
5
6use crate::env::Env;
7use azure_core::{
8    credentials::AccessToken,
9    error::{Error, ErrorKind, Result},
10};
11use std::{
12    ffi::{OsStr, OsString},
13    fmt, io,
14    process::Output,
15    sync::Arc,
16};
17
18mod standard;
19#[cfg(feature = "tokio")]
20mod tokio;
21
22#[allow(unused)]
23pub use standard::StdExecutor;
24#[cfg(feature = "tokio")]
25pub use tokio::TokioExecutor;
26
27/// Creates a new [`Executor`].
28///
29/// The returned Executor spawns a [`std::process::Command`] in a separate thread unless `tokio` is enabled,
30/// in which case it spawns a `tokio::process::Command`.
31pub fn new_executor() -> Arc<dyn Executor> {
32    #[cfg(not(feature = "tokio"))]
33    {
34        Arc::new(StdExecutor)
35    }
36    #[cfg(feature = "tokio")]
37    {
38        Arc::new(TokioExecutor)
39    }
40}
41
42/// An async command runner.
43#[async_trait::async_trait]
44pub trait Executor: Send + Sync + fmt::Debug {
45    /// Run a program with the given arguments until it terminates, returning the output.
46    async fn run(&self, program: &OsStr, args: &[&OsStr]) -> io::Result<Output>;
47}
48
49/// Runs a command in the appropriate platform shell and processes the output
50/// using the specified `OutputProcessor`.
51///
52/// - Windows: Runs `cmd /C {command}` in %SYSTEMROOT%
53/// - Everywhere else: Runs `/bin/sh -c {command}` in /bin
54pub(crate) async fn shell_exec<T: OutputProcessor>(
55    executor: Arc<dyn Executor>,
56    #[cfg_attr(not(windows), allow(unused_variables))] env: &Env,
57    command: &OsStr,
58) -> Result<AccessToken> {
59    let (workdir, program, c_switch) = {
60        #[cfg(windows)]
61        {
62            let system_root = env.var_os("SYSTEMROOT").map_err(|_| {
63                Error::with_message(
64                    ErrorKind::Credential,
65                    "SYSTEMROOT environment variable not set",
66                )
67            })?;
68            (system_root, OsStr::new("cmd"), OsStr::new("/C"))
69        }
70        #[cfg(not(windows))]
71        {
72            (
73                OsString::from("/bin"),
74                OsStr::new("/bin/sh"),
75                OsStr::new("-c"),
76            )
77        }
78    };
79
80    let mut command_string = OsString::from("cd ");
81    command_string.push(workdir);
82    command_string.push(" && ");
83    command_string.push(command);
84    let args = &[c_switch, &command_string];
85
86    let status = executor.run(program, args).await;
87
88    match status {
89        Ok(output) if output.status.success() => {
90            T::deserialize_token(&String::from_utf8_lossy(&output.stdout))
91        }
92        Ok(output) => {
93            let stderr = String::from_utf8_lossy(&output.stderr);
94            let message = if let Some(error_message) = T::get_error_message(&stderr) {
95                error_message
96            } else if output.status.code() == Some(127) || stderr.contains("' is not recognized") {
97                format!("{} not found on PATH", T::tool_name())
98            } else {
99                stderr.to_string()
100            };
101            Err(Error::with_message(ErrorKind::Credential, message))
102        }
103        Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
104            let message = format!("{program:?} wasn't found on PATH");
105            Err(Error::with_error(ErrorKind::Credential, e, message))
106        }
107        Err(e) => {
108            let message = format!("{} error: {e}", e.kind());
109            Err(Error::with_error(ErrorKind::Credential, e, message))
110        }
111    }
112}
113
114pub(crate) trait OutputProcessor: Send + Sized + Sync + 'static {
115    /// Deserialize an AccessToken from stdout
116    fn deserialize_token(stdout: &str) -> Result<AccessToken>;
117
118    /// Optionally convert stderr to a user-friendly error message.
119    /// When this method returns None, the error message will include stderr verbatim.
120    fn get_error_message(stderr: &str) -> Option<String>;
121
122    /// Name of the tool used to get the token e.g. "azd"
123    fn tool_name() -> &'static str;
124}