Skip to main content

smol_workflow_engine/agent_providers/
codex.rs

1use super::common::*;
2use super::types::*;
3use crate::environment::EnvironmentPath;
4use anyhow::{bail, Context};
5use serde_json::{json, Map, Value};
6use std::collections::HashMap;
7use std::path::PathBuf;
8
9#[derive(Debug, Clone)]
10pub struct CodexAgentProviderOptions {
11    pub command: Option<String>,
12    pub subcommand: Vec<String>,
13    pub args: Vec<String>,
14    pub cwd: Option<PathBuf>,
15    pub env: HashMap<String, String>,
16    pub timeout_ms: Option<u64>,
17}
18
19impl Default for CodexAgentProviderOptions {
20    fn default() -> Self {
21        Self {
22            command: None,
23            subcommand: vec!["exec".into()],
24            args: Vec::new(),
25            cwd: None,
26            env: HashMap::new(),
27            timeout_ms: None,
28        }
29    }
30}
31
32#[derive(Debug, Clone, Default)]
33pub struct CodexAgentProvider {
34    options: CodexAgentProviderOptions,
35}
36
37impl CodexAgentProvider {
38    pub fn new(options: CodexAgentProviderOptions) -> Self {
39        Self { options }
40    }
41}
42
43#[async_trait::async_trait]
44impl AgentProvider for CodexAgentProvider {
45    fn name(&self) -> &str {
46        "codex"
47    }
48
49    fn schema_mode(&self) -> AgentProviderSchemaMode {
50        AgentProviderSchemaMode::Builtin
51    }
52
53    fn usage_mode(&self) -> AgentProviderUsageMode {
54        AgentProviderUsageMode::Builtin
55    }
56
57    async fn run(&self, input: AgentProviderRunInput) -> anyhow::Result<AgentProviderResult> {
58        run_codex(input, &self.options).await
59    }
60}
61
62async fn run_codex(
63    input: AgentProviderRunInput,
64    options: &CodexAgentProviderOptions,
65) -> anyhow::Result<AgentProviderResult> {
66    let temp = input.environment.create_temp_dir("smol-wf-codex-").await?;
67    let output_path = join_environment_path(&temp, "last-message.txt");
68    let schema_path = join_environment_path(&temp, "schema.json");
69    let command = options.command.as_deref().unwrap_or("codex");
70    let mut args = Vec::new();
71    args.extend(options.subcommand.clone());
72    if args.is_empty() {
73        args.push("exec".into());
74    }
75    args.extend(options.args.clone());
76    if !args.iter().any(|arg| arg == "--skip-git-repo-check") {
77        args.push("--skip-git-repo-check".into());
78    }
79    if let Some(model) = option_str(&input.options, "model") {
80        args.extend(["--model".into(), model]);
81    }
82    args.extend([
83        "--json".into(),
84        "--output-last-message".into(),
85        output_path.0.clone(),
86    ]);
87    let has_schema = option_schema(&input.options).is_some();
88    if let Some(schema) = option_schema(&input.options) {
89        let schema = to_codex_output_schema(schema);
90        input
91            .environment
92            .write_file(
93                &schema_path,
94                serde_json::to_string_pretty(&schema)?.as_bytes(),
95            )
96            .await?;
97        args.extend(["--output-schema".into(), schema_path.0.clone()]);
98    }
99    args.push("-".into());
100
101    let cwd = input.context.cwd.as_deref().or(options.cwd.as_deref());
102    let (stdout, stderr) = run_command(RunCommandRequest {
103        provider: "Codex",
104        command,
105        args: &args,
106        stdin: Some(&input.prompt),
107        cwd,
108        env: &options.env,
109        timeout_ms: options.timeout_ms,
110        environment: input.environment.as_ref(),
111    })
112    .await?;
113    let events = parse_json_lines(&stdout);
114    let session_id = extract_session_id(&events)
115        .context("Codex provider response did not include a session id")?;
116    let final_message =
117        read_final_message(input.environment.as_ref(), &output_path, &events).await?;
118    let output = if has_schema {
119        parse_structured_output(&final_message)?
120    } else {
121        Value::String(final_message.trim_end().to_string())
122    };
123
124    Ok(AgentProviderResult {
125        output,
126        session_id: Some(session_id),
127        model: extract_model(&Value::Array(events.clone()))
128            .or_else(|| option_model(&input.options)),
129        usage: extract_usage(&events),
130        isolation: None,
131        raw: Some(to_json_value(json!({ "events": events, "stderr": stderr }))),
132    })
133}
134
135fn join_environment_path(base: &EnvironmentPath, child: &str) -> EnvironmentPath {
136    EnvironmentPath(format!("{}/{}", base.as_str().trim_end_matches('/'), child))
137}
138
139async fn read_final_message(
140    environment: &dyn crate::environment::AgentExecutionEnvironment,
141    path: &EnvironmentPath,
142    events: &[Value],
143) -> anyhow::Result<String> {
144    match environment.read_file(path).await {
145        Ok(bytes) => {
146            let message = String::from_utf8_lossy(&bytes).into_owned();
147            if !message.trim().is_empty() {
148                return Ok(message);
149            }
150        }
151        Err(error) => {
152            let not_found = error
153                .chain()
154                .find_map(|cause| cause.downcast_ref::<std::io::Error>())
155                .is_some_and(|error| error.kind() == std::io::ErrorKind::NotFound);
156            if !not_found {
157                bail!("Failed to read codex output file: {error}");
158            }
159        }
160    }
161    if let Some(text) = extract_last_assistant_text(events) {
162        Ok(text)
163    } else {
164        bail!("Codex provider did not return a final assistant message")
165    }
166}
167
168fn to_codex_output_schema(schema: &Value) -> Value {
169    match schema {
170        Value::Array(items) => Value::Array(items.iter().map(to_codex_output_schema).collect()),
171        Value::Object(record) => {
172            let mut output = Map::new();
173            for (key, value) in record {
174                output.insert(key.clone(), to_codex_output_schema(value));
175            }
176            if is_object_schema(&output) {
177                let properties = output
178                    .get("properties")
179                    .and_then(Value::as_object)
180                    .cloned()
181                    .unwrap_or_default();
182                output.insert(
183                    "properties".into(),
184                    to_codex_output_schema(&Value::Object(properties)),
185                );
186                output.insert(
187                    "required".into(),
188                    record
189                        .get("required")
190                        .filter(|v| v.is_array())
191                        .cloned()
192                        .unwrap_or_else(|| json!([])),
193                );
194                output.insert("additionalProperties".into(), Value::Bool(false));
195            }
196            Value::Object(output)
197        }
198        _ => schema.clone(),
199    }
200}
201
202fn is_object_schema(schema: &Map<String, Value>) -> bool {
203    schema.get("type") == Some(&Value::String("object".into())) || schema.contains_key("properties")
204}
205
206fn parse_structured_output(text: &str) -> anyhow::Result<Value> {
207    parse_structured_output_seen(text.trim(), &mut Vec::new())
208}
209
210fn parse_structured_output_seen(text: &str, seen: &mut Vec<String>) -> anyhow::Result<Value> {
211    let trimmed = text.trim();
212    if seen.iter().any(|item| item == trimmed) {
213        bail!("Codex provider did not return valid JSON for schema output");
214    }
215    seen.push(trimmed.to_string());
216
217    if let Ok(parsed) = serde_json::from_str::<Value>(trimmed) {
218        if let Value::String(inner) = parsed {
219            return parse_structured_output_seen(&inner, seen);
220        }
221        return Ok(parsed);
222    }
223
224    if let Some(fenced) = extract_fenced_json(trimmed) {
225        return parse_structured_output_seen(fenced, seen);
226    }
227    if let Some(unescaped) = try_unescape_json_like_text(trimmed) {
228        return parse_structured_output_seen(&unescaped, seen);
229    }
230    if let Some(object_text) = extract_likely_json_text(trimmed) {
231        return parse_structured_output_seen(object_text, seen);
232    }
233    bail!("Codex provider did not return valid JSON for schema output")
234}
235
236fn extract_fenced_json(text: &str) -> Option<&str> {
237    let start = text.find("```")?;
238    let after = &text[start + 3..];
239    let after = after.strip_prefix("json").unwrap_or(after).trim_start();
240    let end = after.find("```")?;
241    Some(after[..end].trim())
242}
243
244fn try_unescape_json_like_text(text: &str) -> Option<String> {
245    if !text.contains("\\n") && !text.contains("\\\"") {
246        return None;
247    }
248    serde_json::from_str::<String>(&format!("\"{text}\""))
249        .ok()
250        .or_else(|| {
251            Some(
252                text.replace("\\n", "\n")
253                    .replace("\\t", "\t")
254                    .replace("\\\"", "\""),
255            )
256        })
257}
258
259fn extract_likely_json_text(text: &str) -> Option<&str> {
260    let object = text.find('{').zip(text.rfind('}')).filter(|(s, e)| e > s);
261    let array = text.find('[').zip(text.rfind(']')).filter(|(s, e)| e > s);
262    object.or(array).map(|(s, e)| &text[s..=e])
263}
264
265fn extract_last_assistant_text(events: &[Value]) -> Option<String> {
266    let mut text = None;
267    for event in events {
268        if let Some(candidate) = extract_assistant_text(event) {
269            text = Some(candidate);
270        }
271    }
272    text
273}
274
275fn extract_assistant_text(value: &Value) -> Option<String> {
276    match value {
277        Value::Array(items) => items.iter().rev().find_map(extract_assistant_text),
278        Value::Object(record) => {
279            let text = extract_text(
280                record
281                    .get("text")
282                    .or_else(|| record.get("output"))
283                    .or_else(|| record.get("message"))
284                    .or_else(|| record.get("content"))?,
285            );
286            if (matches!(
287                record.get("role").and_then(Value::as_str),
288                Some("assistant")
289            ) || matches!(
290                record.get("type").and_then(Value::as_str),
291                Some("assistant_message" | "message")
292            )) && text.is_some()
293            {
294                return text;
295            }
296            for key in [
297                "message",
298                "content",
299                "output",
300                "text",
301                "delta",
302                "part",
303                "parts",
304                "item",
305                "event",
306                "data",
307                "properties",
308            ] {
309                if let Some(candidate) = record.get(key).and_then(extract_assistant_text) {
310                    return Some(candidate);
311                }
312            }
313            None
314        }
315        _ => None,
316    }
317}
318
319fn extract_text(value: &Value) -> Option<String> {
320    match value {
321        Value::String(text) => Some(text.clone()),
322        Value::Array(items) => {
323            let text = items
324                .iter()
325                .map(|item| extract_text(item).unwrap_or_default())
326                .collect::<Vec<_>>()
327                .join("");
328            (!text.is_empty()).then_some(text)
329        }
330        Value::Object(record) => record
331            .get("text")
332            .or_else(|| record.get("content"))
333            .or_else(|| record.get("message"))
334            .or_else(|| record.get("output"))
335            .and_then(extract_text),
336        _ => None,
337    }
338}
339
340fn extract_session_id(events: &[Value]) -> Option<String> {
341    for event in events {
342        if event.get("type").and_then(Value::as_str) == Some("session_meta") {
343            if let Some(id) = get_path(event, &["payload", "id"]).and_then(Value::as_str) {
344                return Some(id.to_string());
345            }
346        }
347        if event.get("type").and_then(Value::as_str) == Some("thread.started") {
348            if let Some(id) = event.get("thread_id").and_then(Value::as_str) {
349                return Some(id.to_string());
350            }
351        }
352        if let Some(id) = event
353            .get("session_id")
354            .or_else(|| event.get("sessionId"))
355            .or_else(|| event.get("sessionID"))
356            .and_then(Value::as_str)
357        {
358            return Some(id.to_string());
359        }
360    }
361    None
362}
363
364fn extract_usage(events: &[Value]) -> Option<AgentUsage> {
365    let mut usage = None;
366    for event in events {
367        if let Some(candidate) = find_first_usage_object(event) {
368            usage = Some(merge_usage_right(usage, normalize_usage(&candidate)));
369        }
370    }
371    usage
372}