File size: 4,850 Bytes
afa0cbf | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 | use std::collections::HashMap;
use std::io;
use std::io::BufRead;
use std::io::IsTerminal;
use std::io::Write;
use anyhow::Context;
use anyhow::Result;
use anyhow::bail;
use codex_app_server_protocol::ToolRequestUserInputAnswer;
use codex_app_server_protocol::ToolRequestUserInputParams;
use codex_app_server_protocol::ToolRequestUserInputResponse;
pub(super) fn prompt_for_answers(
params: &ToolRequestUserInputParams,
) -> Result<ToolRequestUserInputResponse> {
let stdin = io::stdin();
if !stdin.is_terminal() {
bail!("request_user_input requires an interactive stdin terminal");
}
let stdout = io::stdout();
prompt_for_answers_with(&mut stdin.lock(), &mut stdout.lock(), params)
}
fn prompt_for_answers_with(
input: &mut impl BufRead,
output: &mut impl Write,
params: &ToolRequestUserInputParams,
) -> Result<ToolRequestUserInputResponse> {
writeln!(
output,
"\n[request_user_input for thread {}, turn {}]",
params.thread_id, params.turn_id
)?;
if !params.is_blocking {
writeln!(output, "This request is non-blocking.")?;
}
let mut answers = HashMap::new();
for question in ¶ms.questions {
writeln!(output, "\n{}: {}", question.header, question.question)?;
let options = question
.options
.as_deref()
.filter(|options| !options.is_empty());
let answer_values = if let Some(options) = options {
for (index, option) in options.iter().enumerate() {
writeln!(
output,
" {}. {} - {}",
index + 1,
option.label,
option.description
)?;
}
if question.is_other {
writeln!(output, " o. Other (free-form)")?;
}
loop {
if question.is_other {
write!(output, "Choose 1-{} or o: ", options.len())?;
} else {
write!(output, "Choose 1-{}: ", options.len())?;
}
output.flush()?;
let mut line = String::new();
if input
.read_line(&mut line)
.context("failed to read request_user_input selection")?
== 0
{
bail!("stdin closed while waiting for request_user_input selection");
}
let selection = line.trim();
if let Ok(index) = selection.parse::<usize>()
&& let Some(option) = index.checked_sub(1).and_then(|index| options.get(index))
{
break vec![option.label.clone()];
}
if let Some(option) = options
.iter()
.find(|option| option.label.eq_ignore_ascii_case(selection))
{
break vec![option.label.clone()];
}
if question.is_other && selection.eq_ignore_ascii_case("o") {
write!(output, "Other: ")?;
output.flush()?;
line.clear();
if input
.read_line(&mut line)
.context("failed to read request_user_input free-form answer")?
== 0
{
bail!("stdin closed while waiting for request_user_input free-form answer");
}
let answer = line.trim();
if !answer.is_empty() {
break vec![format!("user_note: {answer}")];
}
}
writeln!(output, "Invalid selection; try again.")?;
}
} else {
loop {
write!(output, "Answer: ")?;
output.flush()?;
let mut line = String::new();
if input
.read_line(&mut line)
.context("failed to read request_user_input answer")?
== 0
{
bail!("stdin closed while waiting for request_user_input answer");
}
let answer = line.trim();
if !answer.is_empty() {
break vec![format!("user_note: {answer}")];
}
writeln!(output, "Answer cannot be empty; try again.")?;
}
};
answers.insert(
question.id.clone(),
ToolRequestUserInputAnswer {
answers: answer_values,
},
);
}
Ok(ToolRequestUserInputResponse { answers })
}
#[cfg(test)]
#[path = "request_user_input_tests.rs"]
mod tests;
|