diff --git a/codex-rs/core/src/client.rs b/codex-rs/core/src/client.rs index 7316e90456..79be9c9422 100644 --- a/codex-rs/core/src/client.rs +++ b/codex-rs/core/src/client.rs @@ -40,10 +40,18 @@ use crate::util::backoff; /// When serialized as JSON, this produces a valid "Tool" in the OpenAI /// Responses API. -#[derive(Debug, Serialize)] +#[derive(Debug, Clone, Serialize)] +#[serde(tag = "type")] +enum OpenAiTool { + #[serde(rename = "function")] + Function(ResponsesApiTool), + #[serde(rename = "local_shell")] + LocalShell {}, +} + +#[derive(Debug, Clone, Serialize)] struct ResponsesApiTool { name: &'static str, - r#type: &'static str, // "function" description: &'static str, strict: bool, parameters: JsonSchema, @@ -67,7 +75,7 @@ enum JsonSchema { } /// Tool usage specification -static DEFAULT_TOOLS: LazyLock> = LazyLock::new(|| { +static DEFAULT_TOOLS: LazyLock> = LazyLock::new(|| { let mut properties = BTreeMap::new(); properties.insert( "command".to_string(), @@ -78,9 +86,8 @@ static DEFAULT_TOOLS: LazyLock> = LazyLock::new(|| { properties.insert("workdir".to_string(), JsonSchema::String); properties.insert("timeout".to_string(), JsonSchema::Number); - vec![ResponsesApiTool { + vec![OpenAiTool::Function(ResponsesApiTool { name: "shell", - r#type: "function", description: "Runs a shell command, and returns its output.", strict: false, parameters: JsonSchema::Object { @@ -88,9 +95,12 @@ static DEFAULT_TOOLS: LazyLock> = LazyLock::new(|| { required: &["command"], additional_properties: false, }, - }] + })] }); +static DEFAULT_CODEX_MODEL_TOOLS: LazyLock> = + LazyLock::new(|| vec![OpenAiTool::LocalShell {}]); + #[derive(Clone)] pub struct ModelClient { model: String, @@ -152,8 +162,13 @@ impl ModelClient { } // Assemble tool list: built-in tools + any extra tools from the prompt. - let mut tools_json = Vec::with_capacity(DEFAULT_TOOLS.len() + prompt.extra_tools.len()); - for t in DEFAULT_TOOLS.iter() { + let default_tools = if self.model.starts_with("codex") { + &DEFAULT_TOOLS + } else { + &DEFAULT_CODEX_MODEL_TOOLS + }; + let mut tools_json = Vec::with_capacity(default_tools.len() + prompt.extra_tools.len()); + for t in default_tools.iter() { tools_json.push(serde_json::to_value(t)?); } tools_json.extend( diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index e3cd1a7ad7..301eccb9b2 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -1025,6 +1025,15 @@ async fn handle_response_item( handle_function_call(sess, sub_id.to_string(), name, arguments, call_id).await, ); } + ResponseItem::LocalShellCall { + id, + call_id, + status, + action, + } => { + let _ = (id, call_id, status, action); + todo!() + } ResponseItem::FunctionCallOutput { .. } => { debug!("unexpected FunctionCallOutput from stream"); } diff --git a/codex-rs/core/src/models.rs b/codex-rs/core/src/models.rs index a8817cf7ff..c7ccee1b35 100644 --- a/codex-rs/core/src/models.rs +++ b/codex-rs/core/src/models.rs @@ -1,3 +1,5 @@ +use std::collections::HashMap; + use base64::Engine; use serde::Deserialize; use serde::Serialize; @@ -37,6 +39,14 @@ pub enum ResponseItem { id: String, summary: Vec, }, + LocalShellCall { + /// Set when using the chat completions API. + id: Option, + /// Set when using the Responses API. + call_id: Option, + status: LocalShellStatus, + action: LocalShellAction, + }, FunctionCall { name: String, // The Responses API returns the function call arguments as a *string* that contains @@ -71,6 +81,29 @@ impl From for ResponseItem { } } +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum LocalShellStatus { + Completed, + InProgress, + Incomplete, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum LocalShellAction { + Exec(LocalShellExecAction), +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct LocalShellExecAction { + command: Vec, + timeout_ms: Option, + working_directory: Option, + env: Option>, + user: Option, +} + #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] pub enum ReasoningItemReasoningSummary { diff --git a/codex-rs/core/src/rollout.rs b/codex-rs/core/src/rollout.rs index 4127b603e8..c18a58df06 100644 --- a/codex-rs/core/src/rollout.rs +++ b/codex-rs/core/src/rollout.rs @@ -115,6 +115,7 @@ impl RolloutRecorder { // "fully qualified MCP tool calls," so we could consider // reformatting them in that case. ResponseItem::Message { .. } + | ResponseItem::LocalShellCall { .. } | ResponseItem::FunctionCall { .. } | ResponseItem::FunctionCallOutput { .. } => {} ResponseItem::Reasoning { .. } | ResponseItem::Other => {