Skip to main content

systemprompt_api/routes/admin/
cli.rs

1//! Admin CLI gateway route: streams `systemprompt` subprocess output over SSE.
2//!
3//! Exposes a single authenticated endpoint that validates and forwards an argv
4//! to the CLI binary, propagating the caller's session/context/auth into the
5//! child's environment and relaying stdout/stderr as [`CliOutputEvent`] frames.
6//! The child runs in its own process group and is killed, with its output
7//! forwarders aborted, when the SSE stream is dropped — a client that
8//! disconnects does not leave the command running.
9//!
10//! Copyright (c) systemprompt.io — Business Source License 1.1.
11//! See <https://systemprompt.io> for licensing details.
12
13use axum::extract::Extension;
14use axum::response::IntoResponse;
15use axum::response::sse::{Event, KeepAlive, Sse};
16use axum::routing::post;
17use axum::{Json, Router};
18use futures_util::stream::Stream;
19use std::convert::Infallible;
20use std::sync::Arc;
21use std::time::Duration;
22use systemprompt_events::ToSse;
23use systemprompt_logging::sanitize::redact_argv;
24use systemprompt_models::RequestContext;
25use systemprompt_models::api::{ApiError, CliExecuteRequest, CliOutputEvent};
26use systemprompt_runtime::AppContext;
27use systemprompt_traits::OwnedTask;
28use tokio::io::{AsyncBufReadExt, BufReader};
29use tokio::process::Command;
30
31fn cli_event_to_sse(event: &CliOutputEvent) -> Event {
32    event
33        .to_sse()
34        .unwrap_or_else(|_| Event::default().event("cli").data("{}"))
35}
36
37const MAX_TIMEOUT_SECS: u64 = 600;
38const DEFAULT_CLI_BINARY_PATH: &str = "/app/bin/systemprompt";
39const MAX_CLI_ARGS: usize = 32;
40
41#[derive(Clone, Debug)]
42pub struct CliBinaryPath(Arc<str>);
43
44impl CliBinaryPath {
45    pub fn new(path: impl AsRef<str>) -> Self {
46        Self(Arc::from(path.as_ref()))
47    }
48
49    fn as_str(&self) -> &str {
50        &self.0
51    }
52}
53
54impl Default for CliBinaryPath {
55    fn default() -> Self {
56        Self::new(DEFAULT_CLI_BINARY_PATH)
57    }
58}
59const MAX_CLI_ARG_LEN: usize = 256;
60
61fn validate_cli_args(args: &[String]) -> Result<(), Box<ApiError>> {
62    if args.is_empty() {
63        return Err(Box::new(ApiError::bad_request(
64            "cli args must not be empty",
65        )));
66    }
67    if args.len() > MAX_CLI_ARGS {
68        return Err(Box::new(ApiError::bad_request(format!(
69            "too many cli args (max {MAX_CLI_ARGS})"
70        ))));
71    }
72    for (i, arg) in args.iter().enumerate() {
73        if arg.len() > MAX_CLI_ARG_LEN {
74            return Err(Box::new(ApiError::bad_request(format!(
75                "cli arg #{i} exceeds {MAX_CLI_ARG_LEN} bytes"
76            ))));
77        }
78        if arg
79            .chars()
80            .any(|c| c.is_control() || matches!(c, '`' | '$' | '|' | ';' | '&' | '\n' | '\r'))
81        {
82            return Err(Box::new(ApiError::bad_request(format!(
83                "cli arg #{i} contains forbidden character"
84            ))));
85        }
86    }
87    let first = &args[0];
88    let first_ok = first.chars().next().is_some_and(|c| c.is_ascii_lowercase())
89        && first
90            .chars()
91            .all(|c| c.is_ascii_lowercase() || c.is_ascii_digit() || c == '-');
92    if !first_ok {
93        return Err(Box::new(ApiError::bad_request(
94            "first cli arg must be a lowercase subcommand",
95        )));
96    }
97    Ok(())
98}
99
100pub(super) fn router() -> Router<AppContext> {
101    router_with_binary(CliBinaryPath::default())
102}
103
104pub fn router_with_binary(binary: CliBinaryPath) -> Router<AppContext> {
105    Router::new()
106        .route("/", post(execute_cli))
107        .layer(Extension(binary))
108}
109
110async fn execute_cli(
111    Extension(req_ctx): Extension<RequestContext>,
112    Extension(binary): Extension<CliBinaryPath>,
113    Json(request): Json<CliExecuteRequest>,
114) -> Result<impl IntoResponse, ApiError> {
115    let args = request.args;
116    validate_cli_args(&args).map_err(|e| *e)?;
117    let timeout_secs = request.timeout_secs.min(MAX_TIMEOUT_SECS);
118    let timeout = Duration::from_secs(timeout_secs);
119
120    tracing::info!(
121        user_id = %req_ctx.user_id(),
122        args = ?redact_argv(&args),
123        timeout_secs = timeout_secs,
124        "CLI gateway: executing command"
125    );
126
127    let auth_token = req_ctx.auth_token().map(|token| token.as_str().to_owned());
128    let context_id = req_ctx.context_id().to_string();
129    let session_env = SessionEnv {
130        session: req_ctx.session_id().to_string(),
131        context: context_id,
132        user: req_ctx.user_id().to_string(),
133        auth_token,
134    };
135
136    let stream = create_cli_stream(binary, args, timeout, session_env);
137
138    Ok(Sse::new(stream).keep_alive(KeepAlive::default()))
139}
140
141struct SessionEnv {
142    session: String,
143    context: String,
144    user: String,
145    auth_token: Option<String>,
146}
147
148fn build_cli_command(binary: &CliBinaryPath, args: &[String], session_env: &SessionEnv) -> Command {
149    let mut cmd = Command::new(binary.as_str());
150    cmd.args(args)
151        .env("SYSTEMPROMPT_CLI_REMOTE", "1")
152        .env("SYSTEMPROMPT_SESSION_ID", &session_env.session)
153        .env("SYSTEMPROMPT_CONTEXT_ID", &session_env.context)
154        .env("SYSTEMPROMPT_USER_ID", &session_env.user)
155        .stdout(std::process::Stdio::piped())
156        .stderr(std::process::Stdio::piped())
157        .kill_on_drop(true);
158    systemprompt_loader::subprocess::place_in_own_process_group(cmd.as_std_mut());
159
160    if let Some(token) = &session_env.auth_token {
161        cmd.env("SYSTEMPROMPT_AUTH_TOKEN", token);
162    }
163    cmd
164}
165
166fn spawn_line_forwarder<R>(
167    reader: R,
168    tx: tokio::sync::mpsc::Sender<CliOutputEvent>,
169    make_event: fn(String) -> CliOutputEvent,
170) -> OwnedTask<()>
171where
172    R: tokio::io::AsyncRead + Unpin + Send + 'static,
173{
174    OwnedTask::spawn("cli_output_forwarder", async move {
175        let mut lines = BufReader::new(reader).lines();
176        while let Ok(Some(line)) = lines.next_line().await {
177            if tx.send(make_event(format!("{line}\n"))).await.is_err() {
178                break;
179            }
180        }
181    })
182}
183
184async fn kill_on_timeout(
185    child: &mut tokio::process::Child,
186    timeout: Duration,
187) -> [CliOutputEvent; 2] {
188    tracing::warn!(timeout_secs = timeout.as_secs(), "CLI command timed out");
189    if let Err(e) = child.kill().await {
190        tracing::error!(error = %e, "Failed to kill CLI process");
191    }
192    [
193        CliOutputEvent::Error {
194            message: format!("Timeout after {}s", timeout.as_secs()),
195        },
196        CliOutputEvent::ExitCode { code: -1 },
197    ]
198}
199
200async fn wait_exit_events(mut child: tokio::process::Child) -> Vec<CliOutputEvent> {
201    match child.wait().await {
202        Ok(status) => {
203            let code = status.code().unwrap_or_else(|| {
204                tracing::debug!("Process terminated by signal");
205                -1
206            });
207            tracing::info!(exit_code = code, "CLI command completed");
208            vec![CliOutputEvent::ExitCode { code }]
209        },
210        Err(e) => {
211            tracing::error!(error = %e, "Failed to wait for CLI process");
212            vec![
213                CliOutputEvent::Error {
214                    message: e.to_string(),
215                },
216                CliOutputEvent::ExitCode { code: -1 },
217            ]
218        },
219    }
220}
221
222fn create_cli_stream(
223    binary: CliBinaryPath,
224    args: Vec<String>,
225    timeout: Duration,
226    session_env: SessionEnv,
227) -> impl Stream<Item = Result<Event, Infallible>> {
228    async_stream::stream! {
229        let mut child = match build_cli_command(&binary, &args, &session_env).spawn() {
230            Ok(c) => c,
231            Err(e) => {
232                tracing::error!(error = %e, "Failed to spawn CLI process");
233                yield Ok(cli_event_to_sse(&CliOutputEvent::Error { message: e.to_string() }));
234                yield Ok(cli_event_to_sse(&CliOutputEvent::ExitCode { code: 1 }));
235                return;
236            }
237        };
238
239        let pid = child.id().unwrap_or_else(|| {
240            tracing::debug!("Child process has no PID (already exited?)");
241            0
242        });
243        yield Ok(cli_event_to_sse(&CliOutputEvent::Started { pid }));
244
245        let (tx, mut rx) = tokio::sync::mpsc::channel::<CliOutputEvent>(100);
246        let _stdout_forwarder = child.stdout.take().map(|stdout| {
247            spawn_line_forwarder(stdout, tx.clone(), |data| CliOutputEvent::Stdout { data })
248        });
249        let _stderr_forwarder = child.stderr.take().map(|stderr| {
250            spawn_line_forwarder(stderr, tx.clone(), |data| CliOutputEvent::Stderr { data })
251        });
252        drop(tx);
253
254        let deadline = tokio::time::Instant::now() + timeout;
255
256        loop {
257            tokio::select! {
258                event = rx.recv() => {
259                    match event {
260                        Some(ref e) => yield Ok(cli_event_to_sse(e)),
261                        None => break,
262                    }
263                }
264                () = tokio::time::sleep_until(deadline) => {
265                    for event in kill_on_timeout(&mut child, timeout).await {
266                        yield Ok(cli_event_to_sse(&event));
267                    }
268                    return;
269                }
270            }
271        }
272
273        for event in wait_exit_events(child).await {
274            yield Ok(cli_event_to_sse(&event));
275        }
276    }
277}