Skip to main content

kcode_k1_codex_websearch/
lib.rs

1#![forbid(unsafe_code)]
2#![doc = include_str!("../Documentation.md")]
3#[cfg(not(unix))]
4compile_error!("kcode-k1-codex-websearch requires Unix process groups");
5
6#[cfg(test)]
7use nix::sys::signal::kill;
8use nix::{
9    sys::signal::{Signal, killpg},
10    unistd::Pid,
11};
12use serde_json::Value;
13use std::{
14    ffi::OsString,
15    io,
16    path::PathBuf,
17    time::{Duration, Instant},
18};
19use tokio::{
20    io::{AsyncRead, AsyncReadExt, AsyncWriteExt},
21    process::Command,
22    sync::oneshot,
23};
24
25const ARG_ERROR: &str =
26    "WebSearch failed: model or reasoning_effort is not representable as a process argument";
27const DEADLINE_ERROR: &str = "WebSearch failed: deadline exceeded";
28const COMPAT_ERROR: &str = "WebSearch failed: Codex compatibility check failed";
29const EXEC_ERROR: &str = "WebSearch failed: Codex execution failed";
30const JSON_ERROR: &str = "WebSearch failed: invalid Codex JSONL output";
31const INCOMPLETE_ERROR: &str = "WebSearch failed: Codex response was incomplete";
32const MESSAGE_ERROR: &str = "WebSearch failed: Codex returned no agent message";
33const CLEANUP_RESERVE: Duration = Duration::from_secs(1);
34const HELP_LIMIT: Duration = Duration::from_secs(2);
35
36const CONFIG: &[&str] = &[
37    r#"web_search="live""#,
38    "tools.web_search=true",
39    "tools.view_image=false",
40    "apps._default.enabled=false",
41    "agents.enabled=false",
42    "features.apps=false",
43    "features.code_mode.enabled=false",
44    "features.goals=false",
45    "features.hooks=false",
46    "features.memories=false",
47    "features.multi_agent=false",
48    "features.remote_plugin=false",
49    "features.shell_snapshot=false",
50    "features.shell_tool=false",
51    "features.skill_mcp_dependency_install=false",
52    "features.unified_exec=false",
53    "memories.generate_memories=false",
54    "memories.use_memories=false",
55    r#"history.persistence="none""#,
56    "check_for_update_on_startup=false",
57    "feedback.enabled=false",
58    "analytics.enabled=false",
59    "allow_login_shell=false",
60    "skills.config=[]",
61    "mcp_servers={}",
62    "plugins={}",
63    "marketplaces={}",
64    "hooks={}",
65];
66const REQUIRED_FLAGS: &[&str] = &[
67    "--search",
68    "--ephemeral",
69    "--ignore-user-config",
70    "--ignore-rules",
71    "--json",
72    "--sandbox",
73    "--ask-for-approval",
74    "--skip-git-repo-check",
75    "--no-daemon",
76    "--strict-config",
77    "--model",
78    "--config",
79];
80
81#[derive(Clone)]
82pub struct Runner {
83    executable: PathBuf,
84}
85
86pub struct Request {
87    pub query: String,
88    pub model: String,
89    pub reasoning_effort: String,
90    pub deadline: Instant,
91}
92
93impl Runner {
94    pub fn new(executable: PathBuf) -> Self {
95        Self { executable }
96    }
97
98    pub async fn run(&self, request: Request) -> Result<String, String> {
99        if request.model.contains('\0') || request.reasoning_effort.contains('\0') {
100            return Err(ARG_ERROR.to_owned());
101        }
102        let cutoff = request
103            .deadline
104            .checked_sub(CLEANUP_RESERVE)
105            .unwrap_or(request.deadline);
106        if Instant::now() >= cutoff {
107            return Err(DEADLINE_ERROR.to_owned());
108        }
109
110        let mut help = self.check_help(&["--help"], cutoff).await?;
111        help.extend(self.check_help(&["exec", "--help"], cutoff).await?);
112        let help = String::from_utf8_lossy(&help);
113        if REQUIRED_FLAGS.iter().any(|flag| !help.contains(flag)) {
114            return Err(COMPAT_ERROR.to_owned());
115        }
116        if Instant::now() >= cutoff {
117            return Err(DEADLINE_ERROR.to_owned());
118        }
119
120        let output = run_process(
121            self.executable.clone(),
122            exec_args(&request.model, &request.reasoning_effort),
123            request.query.into_bytes(),
124            cutoff,
125        )
126        .await
127        .map_err(|error| match error {
128            ProcError::Timeout => DEADLINE_ERROR.to_owned(),
129            ProcError::Io | ProcError::Cancelled => EXEC_ERROR.to_owned(),
130        })?;
131        if !output.success {
132            return Err(EXEC_ERROR.to_owned());
133        }
134        parse_jsonl(&output.stdout)
135    }
136
137    async fn check_help(&self, args: &[&str], cutoff: Instant) -> Result<Vec<u8>, String> {
138        if Instant::now() >= cutoff {
139            return Err(DEADLINE_ERROR.to_owned());
140        }
141        let until = cutoff.min(Instant::now() + HELP_LIMIT);
142        let output = run_process(
143            self.executable.clone(),
144            args.iter().map(OsString::from).collect(),
145            Vec::new(),
146            until,
147        )
148        .await
149        .map_err(|error| match error {
150            ProcError::Timeout if Instant::now() >= cutoff => DEADLINE_ERROR.to_owned(),
151            ProcError::Timeout | ProcError::Io | ProcError::Cancelled => COMPAT_ERROR.to_owned(),
152        })?;
153        if !output.success {
154            return Err(COMPAT_ERROR.to_owned());
155        }
156        let mut text = output.stdout;
157        text.extend(output.stderr);
158        Ok(text)
159    }
160}
161
162fn exec_args(model: &str, effort: &str) -> Vec<OsString> {
163    let mut args = [
164        "exec",
165        "--search",
166        "--ephemeral",
167        "--ignore-user-config",
168        "--ignore-rules",
169        "--json",
170        "--sandbox",
171        "read-only",
172        "--ask-for-approval",
173        "never",
174        "--skip-git-repo-check",
175        "--no-daemon",
176        "--strict-config",
177        "--model",
178        model,
179    ]
180    .into_iter()
181    .map(OsString::from)
182    .collect::<Vec<_>>();
183    args.push("-c".into());
184    args.push(format!("model_reasoning_effort={}", toml_quote(effort)).into());
185    for config in CONFIG {
186        args.push("-c".into());
187        args.push((*config).into());
188    }
189    args.push("-".into());
190    args
191}
192
193fn toml_quote(value: &str) -> String {
194    let mut output = String::from("\"");
195    for character in value.chars() {
196        match character {
197            '"' => output.push_str("\\\""),
198            '\\' => output.push_str("\\\\"),
199            '\u{8}' => output.push_str("\\b"),
200            '\t' => output.push_str("\\t"),
201            '\n' => output.push_str("\\n"),
202            '\u{c}' => output.push_str("\\f"),
203            '\r' => output.push_str("\\r"),
204            value if value.is_control() => output.push_str(&format!("\\u{:04X}", value as u32)),
205            value => output.push(value),
206        }
207    }
208    output.push('"');
209    output
210}
211
212fn parse_jsonl(bytes: &[u8]) -> Result<String, String> {
213    let mut turn_completed = false;
214    let mut turn_failed = false;
215    let mut message = None;
216    for raw in bytes.split(|byte| *byte == b'\n') {
217        let line = trim_ascii(raw);
218        if line.is_empty() {
219            continue;
220        }
221        let value: Value = serde_json::from_slice(line).map_err(|_| JSON_ERROR.to_owned())?;
222        match value.get("type").and_then(Value::as_str) {
223            Some("turn.completed") => turn_completed = true,
224            Some("turn.failed") => turn_failed = true,
225            Some("item.completed")
226                if value.pointer("/item/type").and_then(Value::as_str) == Some("agent_message") =>
227            {
228                message = Some(
229                    value
230                        .pointer("/item/text")
231                        .and_then(Value::as_str)
232                        .ok_or_else(|| JSON_ERROR.to_owned())?
233                        .to_owned(),
234                );
235            }
236            _ => {}
237        }
238    }
239    if turn_failed {
240        return Err(EXEC_ERROR.to_owned());
241    }
242    if !turn_completed {
243        return Err(INCOMPLETE_ERROR.to_owned());
244    }
245    message.ok_or_else(|| MESSAGE_ERROR.to_owned())
246}
247
248fn trim_ascii(mut bytes: &[u8]) -> &[u8] {
249    while bytes.first().is_some_and(u8::is_ascii_whitespace) {
250        bytes = &bytes[1..];
251    }
252    while bytes.last().is_some_and(u8::is_ascii_whitespace) {
253        bytes = &bytes[..bytes.len() - 1];
254    }
255    bytes
256}
257
258struct ProcessOutput {
259    success: bool,
260    stdout: Vec<u8>,
261    stderr: Vec<u8>,
262}
263enum ProcError {
264    Io,
265    Timeout,
266    Cancelled,
267}
268struct CancelGuard(Option<oneshot::Sender<()>>);
269impl Drop for CancelGuard {
270    fn drop(&mut self) {
271        if let Some(sender) = self.0.take() {
272            let _ = sender.send(());
273        }
274    }
275}
276
277async fn drain<R: AsyncRead + Unpin>(mut reader: R) -> io::Result<Vec<u8>> {
278    let mut bytes = Vec::new();
279    reader.read_to_end(&mut bytes).await?;
280    Ok(bytes)
281}
282
283async fn run_process(
284    executable: PathBuf,
285    args: Vec<OsString>,
286    input: Vec<u8>,
287    until: Instant,
288) -> Result<ProcessOutput, ProcError> {
289    if Instant::now() >= until {
290        return Err(ProcError::Timeout);
291    }
292    let directory = tempfile::tempdir().map_err(|_| ProcError::Io)?;
293    let mut command = Command::new(executable);
294    command
295        .args(args)
296        .current_dir(directory.path())
297        .stdin(std::process::Stdio::piped())
298        .stdout(std::process::Stdio::piped())
299        .stderr(std::process::Stdio::piped())
300        .kill_on_drop(true)
301        .process_group(0);
302    let mut child = command.spawn().map_err(|_| ProcError::Io)?;
303    let pid = Pid::from_raw(child.id().ok_or(ProcError::Io)? as i32);
304    let mut stdin = child.stdin.take().ok_or(ProcError::Io)?;
305    let stdout = child.stdout.take().ok_or(ProcError::Io)?;
306    let stderr = child.stderr.take().ok_or(ProcError::Io)?;
307    let (cancel_tx, mut cancel_rx) = oneshot::channel();
308
309    let worker = tokio::spawn(async move {
310        let _directory = directory;
311        let complete = async move {
312            let write = async move {
313                stdin.write_all(&input).await?;
314                stdin.shutdown().await
315            };
316            let (status, written, stdout, stderr) =
317                tokio::join!(child.wait(), write, drain(stdout), drain(stderr));
318            written.map_err(|_| ProcError::Io)?;
319            Ok::<_, ProcError>(ProcessOutput {
320                success: status.map_err(|_| ProcError::Io)?.success(),
321                stdout: stdout.map_err(|_| ProcError::Io)?,
322                stderr: stderr.map_err(|_| ProcError::Io)?,
323            })
324        };
325        tokio::pin!(complete);
326        let timer = tokio::time::sleep_until(tokio::time::Instant::from_std(until));
327        tokio::pin!(timer);
328        tokio::select! {
329            result = &mut complete => result,
330            _ = &mut timer => {
331                let _ = killpg(pid, Signal::SIGKILL);
332                let _ = (&mut complete).await;
333                Err(ProcError::Timeout)
334            }
335            _ = &mut cancel_rx => {
336                let _ = killpg(pid, Signal::SIGKILL);
337                let _ = (&mut complete).await;
338                Err(ProcError::Cancelled)
339            }
340        }
341    });
342    let mut guard = CancelGuard(Some(cancel_tx));
343    let result = worker.await.map_err(|_| ProcError::Io)?;
344    guard.0.take();
345    result
346}
347
348#[cfg(test)]
349mod tests {
350    use super::*;
351    use std::{fs, os::unix::fs::PermissionsExt, path::Path};
352
353    struct Fake {
354        _directory: tempfile::TempDir,
355        executable: PathBuf,
356        args: PathBuf,
357        input: PathBuf,
358        pids: PathBuf,
359    }
360    fn quote(path: &Path) -> String {
361        format!("'{}'", path.display().to_string().replace('\'', "'\"'\"'"))
362    }
363    fn fake(body: &str, exit: i32) -> Fake {
364        let directory = tempfile::tempdir().unwrap();
365        let executable = directory.path().join("codex");
366        let args = directory.path().join("args");
367        let input = directory.path().join("input");
368        let pids = directory.path().join("pids");
369        let body = body.replace("@PIDS@", &quote(&pids));
370        let script = format!(
371            "#!/bin/sh\nif [ \"$1\" = \"--help\" ] || {{ [ \"$1\" = exec ] && [ \"$2\" = \"--help\" ]; }}; then\nprintf '%s\\n' '{}'\nexit 0\nfi\nprintf '%s\\n' \"$@\" > {}\ncat > {}\n{}\nexit {}\n",
372            REQUIRED_FLAGS.join(" "),
373            quote(&args),
374            quote(&input),
375            body,
376            exit,
377        );
378        fs::write(&executable, script).unwrap();
379        fs::set_permissions(&executable, fs::Permissions::from_mode(0o755)).unwrap();
380        Fake {
381            _directory: directory,
382            executable,
383            args,
384            input,
385            pids,
386        }
387    }
388    fn request(query: &str, duration: Duration) -> Request {
389        Request {
390            query: query.into(),
391            model: "model-x".into(),
392            reasoning_effort: "high".into(),
393            deadline: Instant::now() + duration,
394        }
395    }
396
397    #[tokio::test]
398    async fn invocation_stdin_and_last_message_are_exact() {
399        let fake = fake(
400            r#"printf '%s\n' \
401'{"type":"future.event","x":1}' \
402'{"type":"item.completed","item":{"type":"agent_message","text":"old"}}' \
403'{"type":"item.completed","item":{"type":"agent_message","text":"最後\nline"}}' \
404'{"type":"turn.completed","extra":true}'"#,
405            0,
406        );
407        let query = "héllo\n世界\0tail";
408        let answer = Runner::new(fake.executable.clone())
409            .run(request(query, Duration::from_secs(5)))
410            .await
411            .unwrap();
412        assert_eq!(answer, "最後\nline");
413        assert_eq!(fs::read(&fake.input).unwrap(), query.as_bytes());
414        let text = fs::read_to_string(&fake.args).unwrap();
415        let got = text.lines().collect::<Vec<_>>();
416        assert_eq!(
417            &got[..15],
418            [
419                "exec",
420                "--search",
421                "--ephemeral",
422                "--ignore-user-config",
423                "--ignore-rules",
424                "--json",
425                "--sandbox",
426                "read-only",
427                "--ask-for-approval",
428                "never",
429                "--skip-git-repo-check",
430                "--no-daemon",
431                "--strict-config",
432                "--model",
433                "model-x"
434            ]
435        );
436        assert_eq!(&got[15..17], ["-c", r#"model_reasoning_effort="high""#]);
437        for (index, config) in CONFIG.iter().enumerate() {
438            assert_eq!(&got[17 + index * 2..19 + index * 2], ["-c", *config]);
439        }
440        assert_eq!(got.last(), Some(&"-"));
441        assert_eq!(toml_quote("a\"\n\u{7f}"), "\"a\\\"\\n\\u007F\"");
442    }
443
444    #[tokio::test]
445    async fn rejects_failures_and_nul_arguments() {
446        let cases = [
447            ("printf '%s\\n' '{'", 0),
448            (r#"printf '%s\n' '{"type":"turn.failed"}'"#, 0),
449            (r#"printf '%s\n' '{"type":"turn.completed"}'"#, 7),
450            (r#"printf '%s\n' '{"type":"turn.completed"}'"#, 0),
451            (
452                r#"printf '%s\n' '{"type":"item.completed","item":{"type":"agent_message","text":"x"}}'"#,
453                0,
454            ),
455        ];
456        for (body, status) in cases {
457            let fake = fake(body, status);
458            assert!(
459                Runner::new(fake.executable)
460                    .run(request("q", Duration::from_secs(5)))
461                    .await
462                    .is_err()
463            );
464        }
465        for model in [true, false] {
466            let runner = Runner::new("/not/executed".into());
467            let mut request = request("q", Duration::from_secs(2));
468            if model {
469                request.model.push('\0');
470            } else {
471                request.reasoning_effort.push('\0');
472            }
473            assert_eq!(runner.run(request).await.unwrap_err(), ARG_ERROR);
474        }
475    }
476
477    async fn wait_for_file(path: &Path) {
478        for _ in 0..200 {
479            if path.exists() {
480                return;
481            }
482            tokio::time::sleep(Duration::from_millis(10)).await;
483        }
484        panic!("fake process did not start");
485    }
486    fn process_is_live(pid: i32) -> bool {
487        if kill(Pid::from_raw(pid), None).is_err() {
488            return false;
489        }
490        let Ok(stat) = fs::read_to_string(format!("/proc/{pid}/stat")) else {
491            return true;
492        };
493        let Some(tail) = stat.rsplit_once(") ").map(|(_, tail)| tail) else {
494            return true;
495        };
496        !matches!(tail.as_bytes().first(), Some(b'Z' | b'X'))
497    }
498    async fn wait_for_processes_to_stop(path: &Path) {
499        let pids = fs::read_to_string(path)
500            .unwrap()
501            .split_whitespace()
502            .map(|value| value.parse::<i32>().unwrap())
503            .collect::<Vec<_>>();
504        assert!(pids.len() >= 2);
505        for _ in 0..200 {
506            if pids.iter().all(|pid| !process_is_live(*pid)) {
507                return;
508            }
509            tokio::time::sleep(Duration::from_millis(10)).await;
510        }
511        panic!("process-group member survived");
512    }
513
514    #[tokio::test]
515    async fn deadline_and_cancellation_kill_descendants() {
516        let deadline = fake(r#"sleep 30 & echo "$$ $!" > @PIDS@; wait"#, 0);
517        let error = Runner::new(deadline.executable.clone())
518            .run(request("q", Duration::from_millis(1600)))
519            .await
520            .unwrap_err();
521        assert_eq!(error, DEADLINE_ERROR);
522        wait_for_processes_to_stop(&deadline.pids).await;
523
524        let cancelled = fake(r#"sleep 30 & echo "$$ $!" > @PIDS@; wait"#, 0);
525        let runner = Runner::new(cancelled.executable.clone());
526        let task =
527            tokio::spawn(async move { runner.run(request("q", Duration::from_secs(10))).await });
528        wait_for_file(&cancelled.pids).await;
529        task.abort();
530        let _ = task.await;
531        wait_for_processes_to_stop(&cancelled.pids).await;
532    }
533}