fix: keep CLI provider prompts out of process arguments (#10547)
This commit is contained in:
@@ -4,7 +4,7 @@ use rmcp::model::Role;
|
||||
use serde_json::{json, Value};
|
||||
use std::path::PathBuf;
|
||||
use std::process::Stdio;
|
||||
use tokio::io::{AsyncBufReadExt, BufReader};
|
||||
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
|
||||
use tokio::process::Command;
|
||||
|
||||
use super::base::{
|
||||
@@ -208,15 +208,17 @@ impl CursorAgentProvider {
|
||||
|
||||
cmd.arg("--model").arg(&model.model_name);
|
||||
|
||||
cmd.arg("-p")
|
||||
.arg(&prompt)
|
||||
cmd.arg("--print")
|
||||
.arg("--output-format")
|
||||
.arg("json")
|
||||
.arg("--force");
|
||||
|
||||
cmd.stdout(Stdio::piped()).stderr(Stdio::piped());
|
||||
cmd.stdin(Stdio::piped())
|
||||
.stdout(Stdio::piped())
|
||||
.stderr(Stdio::piped());
|
||||
|
||||
let mut child = cmd
|
||||
.kill_on_drop(true)
|
||||
.spawn()
|
||||
.map_err(|e| ProviderError::RequestFailed(format!(
|
||||
"Failed to spawn cursor-agent CLI command '{:?}': {}. \
|
||||
@@ -224,10 +226,27 @@ impl CursorAgentProvider {
|
||||
self.command, e
|
||||
)))?;
|
||||
|
||||
let mut stdin = child
|
||||
.stdin
|
||||
.take()
|
||||
.ok_or_else(|| ProviderError::RequestFailed("Failed to capture stdin".to_string()))?;
|
||||
let prompt_write = tokio::spawn(async move {
|
||||
stdin.write_all(prompt.as_bytes()).await?;
|
||||
stdin.shutdown().await
|
||||
});
|
||||
|
||||
let stdout = child
|
||||
.stdout
|
||||
.take()
|
||||
.ok_or_else(|| ProviderError::RequestFailed("Failed to capture stdout".to_string()))?;
|
||||
let stderr = child.stderr.take();
|
||||
let stderr_drain = tokio::spawn(async move {
|
||||
let mut output = String::new();
|
||||
if let Some(mut stderr) = stderr {
|
||||
let _ = stderr.read_to_string(&mut output).await;
|
||||
}
|
||||
output
|
||||
});
|
||||
|
||||
let mut reader = BufReader::new(stdout);
|
||||
let mut lines = Vec::new();
|
||||
@@ -255,6 +274,8 @@ impl CursorAgentProvider {
|
||||
let exit_status = child.wait().await.map_err(|e| {
|
||||
ProviderError::RequestFailed(format!("Failed to wait for command: {}", e))
|
||||
})?;
|
||||
let prompt_write_result = prompt_write.await;
|
||||
let _stderr = stderr_drain.await.unwrap_or_default();
|
||||
|
||||
if !exit_status.success() {
|
||||
if !self.get_authentication_status().await {
|
||||
@@ -268,6 +289,14 @@ impl CursorAgentProvider {
|
||||
)));
|
||||
}
|
||||
|
||||
prompt_write_result
|
||||
.map_err(|e| {
|
||||
ProviderError::RequestFailed(format!("Failed to write prompt to stdin: {e}"))
|
||||
})?
|
||||
.map_err(|e| {
|
||||
ProviderError::RequestFailed(format!("Failed to write prompt to stdin: {e}"))
|
||||
})?;
|
||||
|
||||
tracing::debug!("Command executed successfully, got {} lines", lines.len());
|
||||
for (i, line) in lines.iter().enumerate() {
|
||||
tracing::debug!("Line {}: {}", i, line);
|
||||
@@ -362,3 +391,73 @@ impl Provider for CursorAgentProvider {
|
||||
Ok(stream_from_single_message(message, provider_usage))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(all(test, unix))]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::fs;
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
use std::path::Path;
|
||||
|
||||
const SENTINEL: &str = "loupe-sensitive-cursor-prompt";
|
||||
|
||||
fn recording_cli(directory: &Path) -> PathBuf {
|
||||
let command = directory.join("cursor-agent-recording-shim");
|
||||
fs::write(
|
||||
&command,
|
||||
r#"#!/bin/sh
|
||||
record_dir=${0%/*}
|
||||
printf '%s\n' "$@" > "$record_dir/args"
|
||||
cat > "$record_dir/stdin"
|
||||
printf '%s\n' '{"type":"result","result":"ok"}'
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
fs::set_permissions(&command, fs::Permissions::from_mode(0o755)).unwrap();
|
||||
command
|
||||
}
|
||||
|
||||
async fn assert_prompt_uses_stdin(messages: Vec<Message>) {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let provider = CursorAgentProvider {
|
||||
command: recording_cli(directory.path()),
|
||||
name: CURSOR_AGENT_PROVIDER_NAME.to_string(),
|
||||
};
|
||||
|
||||
let lines = provider
|
||||
.execute_command(
|
||||
&ModelConfig::new(CURSOR_AGENT_DEFAULT_MODEL),
|
||||
"system instructions",
|
||||
&messages,
|
||||
&[],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(lines, vec![r#"{"type":"result","result":"ok"}"#]);
|
||||
let args = fs::read_to_string(directory.path().join("args")).unwrap();
|
||||
let stdin = fs::read_to_string(directory.path().join("stdin")).unwrap();
|
||||
assert!(!args.contains(SENTINEL));
|
||||
assert!(stdin.contains(SENTINEL));
|
||||
assert!(!args.lines().any(|arg| arg == "-p"));
|
||||
assert!(args.contains("--model\nauto"));
|
||||
assert!(args.lines().any(|arg| arg == "--print"));
|
||||
assert!(args.contains("--output-format\njson"));
|
||||
assert!(args.contains("--force"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn initial_prompt_is_sent_on_stdin() {
|
||||
assert_prompt_uses_stdin(vec![Message::user().with_text(SENTINEL)]).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resumed_conversation_is_sent_on_stdin() {
|
||||
assert_prompt_uses_stdin(vec![
|
||||
Message::user().with_text("first turn"),
|
||||
Message::assistant().with_text("first response"),
|
||||
Message::user().with_text(SENTINEL),
|
||||
])
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,7 +4,7 @@ use serde_json::Value;
|
||||
use std::path::PathBuf;
|
||||
use std::process::Stdio;
|
||||
use std::sync::{Arc, OnceLock};
|
||||
use tokio::io::{AsyncBufReadExt, AsyncReadExt, BufReader};
|
||||
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
|
||||
use tokio::process::Command;
|
||||
|
||||
use super::base::{
|
||||
@@ -44,6 +44,12 @@ pub struct GeminiCliProvider {
|
||||
cli_session_id: Arc<OnceLock<String>>,
|
||||
}
|
||||
|
||||
struct GeminiCliProcess {
|
||||
child: tokio::process::Child,
|
||||
reader: BufReader<tokio::process::ChildStdout>,
|
||||
prompt_write: tokio::task::JoinHandle<std::io::Result<()>>,
|
||||
}
|
||||
|
||||
impl GeminiCliProvider {
|
||||
pub async fn from_env(
|
||||
_tls_config: Option<crate::providers::api_client::TlsConfig>,
|
||||
@@ -91,7 +97,7 @@ impl GeminiCliProvider {
|
||||
}
|
||||
}
|
||||
|
||||
fn build_command(&self, prompt: &str, model_name: &str) -> Command {
|
||||
fn build_command(&self, model_name: &str) -> Command {
|
||||
let mut cmd = Command::new(&self.command);
|
||||
configure_subprocess(&mut cmd);
|
||||
|
||||
@@ -105,13 +111,9 @@ impl GeminiCliProvider {
|
||||
cmd.arg("-r").arg(sid);
|
||||
}
|
||||
|
||||
cmd.arg("-p")
|
||||
.arg(prompt)
|
||||
.arg("--output-format")
|
||||
.arg("stream-json")
|
||||
.arg("--yolo");
|
||||
cmd.arg("--output-format").arg("stream-json").arg("--yolo");
|
||||
|
||||
cmd.stdin(Stdio::null())
|
||||
cmd.stdin(Stdio::piped())
|
||||
.stdout(Stdio::piped())
|
||||
.stderr(Stdio::piped());
|
||||
|
||||
@@ -123,18 +125,12 @@ impl GeminiCliProvider {
|
||||
system: &str,
|
||||
messages: &[Message],
|
||||
model_name: &str,
|
||||
) -> Result<
|
||||
(
|
||||
tokio::process::Child,
|
||||
BufReader<tokio::process::ChildStdout>,
|
||||
),
|
||||
ProviderError,
|
||||
> {
|
||||
) -> Result<GeminiCliProcess, ProviderError> {
|
||||
let prompt = self.build_prompt(system, messages);
|
||||
|
||||
tracing::debug!(command = ?self.command, "Executing Gemini CLI command");
|
||||
|
||||
let mut cmd = self.build_command(&prompt, model_name);
|
||||
let mut cmd = self.build_command(model_name);
|
||||
|
||||
let mut child = cmd.kill_on_drop(true).spawn().map_err(|e| {
|
||||
ProviderError::RequestFailed(format!(
|
||||
@@ -144,12 +140,25 @@ impl GeminiCliProvider {
|
||||
))
|
||||
})?;
|
||||
|
||||
let mut stdin = child
|
||||
.stdin
|
||||
.take()
|
||||
.ok_or_else(|| ProviderError::RequestFailed("Failed to capture stdin".to_string()))?;
|
||||
let prompt_write = tokio::spawn(async move {
|
||||
stdin.write_all(prompt.as_bytes()).await?;
|
||||
stdin.shutdown().await
|
||||
});
|
||||
|
||||
let stdout = child
|
||||
.stdout
|
||||
.take()
|
||||
.ok_or_else(|| ProviderError::RequestFailed("Failed to capture stdout".to_string()))?;
|
||||
|
||||
Ok((child, BufReader::new(stdout)))
|
||||
Ok(GeminiCliProcess {
|
||||
child,
|
||||
reader: BufReader::new(stdout),
|
||||
prompt_write,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -216,8 +225,11 @@ impl Provider for GeminiCliProvider {
|
||||
return Ok(stream_from_single_message(message, provider_usage));
|
||||
}
|
||||
|
||||
let (mut child, mut reader) =
|
||||
self.spawn_command(system, messages, &model_config.model_name)?;
|
||||
let GeminiCliProcess {
|
||||
mut child,
|
||||
mut reader,
|
||||
prompt_write,
|
||||
} = self.spawn_command(system, messages, &model_config.model_name)?;
|
||||
let session_id_lock = Arc::clone(&self.cli_session_id);
|
||||
let model_name = model_config.model_name.clone();
|
||||
let message_id = uuid::Uuid::new_v4().to_string();
|
||||
@@ -297,6 +309,7 @@ impl Provider for GeminiCliProvider {
|
||||
}
|
||||
}
|
||||
|
||||
let prompt_write_result = prompt_write.await;
|
||||
let stderr_text = stderr_drain.await.unwrap_or_default();
|
||||
let exit_status = child.wait().await.map_err(|e| {
|
||||
ProviderError::RequestFailed(format!("Failed to wait for command: {e}"))
|
||||
@@ -314,6 +327,14 @@ impl Provider for GeminiCliProvider {
|
||||
)))?;
|
||||
}
|
||||
|
||||
prompt_write_result
|
||||
.map_err(|e| ProviderError::RequestFailed(format!(
|
||||
"Failed to write prompt to stdin: {e}"
|
||||
)))?
|
||||
.map_err(|e| ProviderError::RequestFailed(format!(
|
||||
"Failed to write prompt to stdin: {e}"
|
||||
)))?;
|
||||
|
||||
let provider_usage = ProviderUsage::new(model_name, accumulated_usage);
|
||||
yield (None, Some(provider_usage));
|
||||
}))
|
||||
@@ -323,6 +344,17 @@ impl Provider for GeminiCliProvider {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[cfg(unix)]
|
||||
use futures::StreamExt;
|
||||
#[cfg(unix)]
|
||||
use std::fs;
|
||||
#[cfg(unix)]
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
#[cfg(unix)]
|
||||
use std::path::Path;
|
||||
|
||||
#[cfg(unix)]
|
||||
const SENTINEL: &str = "loupe-sensitive-gemini-prompt";
|
||||
|
||||
fn make_provider() -> GeminiCliProvider {
|
||||
GeminiCliProvider {
|
||||
@@ -358,4 +390,83 @@ mod tests {
|
||||
let prompt = provider.build_prompt("You are helpful.", &messages);
|
||||
assert_eq!(prompt, "Follow up question");
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
fn recording_cli(directory: &Path) -> PathBuf {
|
||||
let command = directory.join("gemini-recording-shim");
|
||||
fs::write(
|
||||
&command,
|
||||
r#"#!/bin/sh
|
||||
record_dir=${0%/*}
|
||||
printf '%s\n' "$@" > "$record_dir/args"
|
||||
cat > "$record_dir/stdin"
|
||||
printf '%s\n' '{"type":"init","session_id":"recorded-session"}'
|
||||
printf '%s\n' '{"type":"result","stats":{}}'
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
fs::set_permissions(&command, fs::Permissions::from_mode(0o755)).unwrap();
|
||||
command
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
async fn assert_prompt_uses_stdin(resumed: bool) {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let mut provider = make_provider();
|
||||
provider.command = recording_cli(directory.path());
|
||||
if resumed {
|
||||
provider
|
||||
.cli_session_id
|
||||
.set("existing-session".to_string())
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
let messages = if resumed {
|
||||
vec![
|
||||
Message::user().with_text("first turn"),
|
||||
Message::assistant().with_text("first response"),
|
||||
Message::user().with_text(SENTINEL),
|
||||
]
|
||||
} else {
|
||||
vec![Message::user().with_text(SENTINEL)]
|
||||
};
|
||||
let mut stream = provider
|
||||
.stream(
|
||||
&ModelConfig::new(GEMINI_CLI_DEFAULT_MODEL),
|
||||
"system instructions",
|
||||
&messages,
|
||||
&[],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
while let Some(item) = stream.next().await {
|
||||
item.unwrap();
|
||||
}
|
||||
|
||||
let args = fs::read_to_string(directory.path().join("args")).unwrap();
|
||||
let stdin = fs::read_to_string(directory.path().join("stdin")).unwrap();
|
||||
assert!(!args.contains(SENTINEL));
|
||||
assert!(stdin.contains(SENTINEL));
|
||||
assert!(!args.lines().any(|arg| arg == "-p"));
|
||||
assert!(args.contains("-m\ngemini-2.5-pro"));
|
||||
assert!(args.contains("--output-format\nstream-json"));
|
||||
assert!(args.contains("--yolo"));
|
||||
if resumed {
|
||||
assert!(args.contains("-r\nexisting-session"));
|
||||
} else {
|
||||
assert!(!args.lines().any(|arg| arg == "-r"));
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[tokio::test]
|
||||
async fn initial_prompt_is_sent_on_stdin() {
|
||||
assert_prompt_uses_stdin(false).await;
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[tokio::test]
|
||||
async fn resumed_prompt_is_sent_on_stdin() {
|
||||
assert_prompt_uses_stdin(true).await;
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user