From ab33ddb9a613a6055064a45fc7b8e1ec877ef4f7 Mon Sep 17 00:00:00 2001 From: Michael Bolin Date: Wed, 7 May 2025 23:13:52 -0700 Subject: [PATCH] feat: support the chat completions API in the Rust CLI --- codex-rs/core/src/chat_completions.rs | 191 ++++++++++++++++++++ codex-rs/core/src/client.rs | 94 +++------- codex-rs/core/src/client_common.rs | 72 ++++++++ codex-rs/core/src/codex.rs | 4 +- codex-rs/core/src/error.rs | 2 +- codex-rs/core/src/lib.rs | 3 + codex-rs/core/src/model_provider_info.rs | 35 +++- codex-rs/core/tests/previous_response_id.rs | 1 + codex-rs/core/tests/stream_no_completed.rs | 1 + 9 files changed, 326 insertions(+), 77 deletions(-) create mode 100644 codex-rs/core/src/chat_completions.rs create mode 100644 codex-rs/core/src/client_common.rs diff --git a/codex-rs/core/src/chat_completions.rs b/codex-rs/core/src/chat_completions.rs new file mode 100644 index 0000000000..839ec906db --- /dev/null +++ b/codex-rs/core/src/chat_completions.rs @@ -0,0 +1,191 @@ +use std::time::Duration; + +use bytes::Bytes; +use eventsource_stream::Eventsource; +use futures::Stream; +use futures::StreamExt; +use futures::TryStreamExt; +use reqwest::StatusCode; +use serde_json::json; +use tokio::sync::mpsc; +use tokio::time::timeout; +use tracing::debug; +use tracing::trace; + +use crate::ModelProviderInfo; +use crate::client_common::Prompt; +use crate::client_common::ResponseEvent; +use crate::client_common::ResponseStream; +use crate::error::CodexErr; +use crate::error::Result; +use crate::flags::OPENAI_REQUEST_MAX_RETRIES; +use crate::flags::OPENAI_STREAM_IDLE_TIMEOUT_MS; +use crate::models::ContentItem; +use crate::models::ResponseItem; +use crate::util::backoff; + +/// Implementation for the classic Chat Completions API. This is intentionally +/// minimal: we only stream back plain assistant text. +pub(crate) async fn stream_chat_completions( + prompt: &Prompt, + model: &str, + client: &reqwest::Client, + provider: &ModelProviderInfo, +) -> Result { + // Build messages array + let mut messages = Vec::::new(); + + if let Some(instr) = &prompt.instructions { + messages.push(json!({"role": "system", "content": instr})); + } + + for item in &prompt.input { + if let ResponseItem::Message { role, content } = item { + let mut text = String::new(); + for c in content { + match c { + ContentItem::InputText { text: t } | ContentItem::OutputText { text: t } => { + text.push_str(t); + } + _ => {} + } + } + messages.push(json!({"role": role, "content": text})); + } + } + + let payload = json!({ + "model": model, + "messages": messages, + "stream": true + }); + + let base_url = provider.base_url.trim_end_matches('/'); + let url = format!("{}/chat/completions", base_url); + + debug!(url, "POST (chat)"); + trace!("request payload: {}", payload); + + let api_key = provider.api_key()?; + let mut attempt = 0; + loop { + attempt += 1; + + let res = client + .post(&url) + .bearer_auth(api_key.clone()) + .header(reqwest::header::ACCEPT, "text/event-stream") + .json(&payload) + .send() + .await; + + match res { + Ok(resp) if resp.status().is_success() => { + let (tx_event, rx_event) = mpsc::channel::>(16); + let stream = resp.bytes_stream().map_err(CodexErr::Reqwest); + tokio::spawn(process_chat_sse(stream, tx_event)); + return Ok(ResponseStream { rx_event }); + } + Ok(res) => { + let status = res.status(); + if !(status == StatusCode::TOO_MANY_REQUESTS || status.is_server_error()) { + let body = (res.text().await).unwrap_or_default(); + return Err(CodexErr::UnexpectedStatus(status, body)); + } + + if attempt > *OPENAI_REQUEST_MAX_RETRIES { + return Err(CodexErr::RetryLimit(status)); + } + + let retry_after_secs = res + .headers() + .get(reqwest::header::RETRY_AFTER) + .and_then(|v| v.to_str().ok()) + .and_then(|s| s.parse::().ok()); + + let delay = retry_after_secs + .map(|s| Duration::from_millis(s * 1_000)) + .unwrap_or_else(|| backoff(attempt)); + tokio::time::sleep(delay).await; + } + Err(e) => { + if attempt > *OPENAI_REQUEST_MAX_RETRIES { + return Err(e.into()); + } + let delay = backoff(attempt); + tokio::time::sleep(delay).await; + } + } + } +} + +/// Lightweight SSE processor for the Chat Completions streaming format. The +/// output is mapped onto Codex's internal [`ResponseEvent`] so that the rest +/// of the pipeline can stay agnostic of the underlying wire format. +async fn process_chat_sse(stream: S, tx_event: mpsc::Sender>) +where + S: Stream> + Unpin, +{ + let mut stream = stream.eventsource(); + + let idle_timeout = *OPENAI_STREAM_IDLE_TIMEOUT_MS; + + loop { + let sse = match timeout(idle_timeout, stream.next()).await { + Ok(Some(Ok(ev))) => ev, + Ok(Some(Err(e))) => { + let _ = tx_event.send(Err(CodexErr::Stream(e.to_string()))).await; + return; + } + Ok(None) => { + // Stream closed gracefully – emit Completed with dummy id. + let _ = tx_event + .send(Ok(ResponseEvent::Completed { + response_id: String::new(), + })) + .await; + return; + } + Err(_) => { + let _ = tx_event + .send(Err(CodexErr::Stream("idle timeout waiting for SSE".into()))) + .await; + return; + } + }; + + // OpenAI Chat streaming sends a literal string "[DONE]" when finished. + if sse.data.trim() == "[DONE]" { + let _ = tx_event + .send(Ok(ResponseEvent::Completed { + response_id: String::new(), + })) + .await; + return; + } + + // Parse JSON chunk + let chunk: serde_json::Value = match serde_json::from_str(&sse.data) { + Ok(v) => v, + Err(_) => continue, + }; + + let content_opt = chunk + .get("choices") + .and_then(|c| c.get(0)) + .and_then(|c| c.get("delta")) + .and_then(|d| d.get("content")) + .and_then(|c| c.as_str()); + + if let Some(content) = content_opt { + let item = ResponseItem::Message { + role: "assistant".to_string(), + content: vec![ContentItem::OutputText { + text: content.to_string(), + }], + }; + + let _ = tx_event.send(Ok(ResponseEvent::OutputItemDone(item))).await; + } + } +} diff --git a/codex-rs/core/src/client.rs b/codex-rs/core/src/client.rs index 9216e68ce6..9891e68b83 100644 --- a/codex-rs/core/src/client.rs +++ b/codex-rs/core/src/client.rs @@ -1,11 +1,7 @@ use std::collections::BTreeMap; -use std::collections::HashMap; use std::io::BufRead; use std::path::Path; -use std::pin::Pin; use std::sync::LazyLock; -use std::task::Context; -use std::task::Poll; use std::time::Duration; use bytes::Bytes; @@ -23,66 +19,22 @@ use tracing::debug; use tracing::trace; use tracing::warn; +use crate::chat_completions::stream_chat_completions; +use crate::client_common::Payload; +use crate::client_common::Prompt; +use crate::client_common::Reasoning; +use crate::client_common::ResponseEvent; +use crate::client_common::ResponseStream; use crate::error::CodexErr; use crate::error::Result; use crate::flags::CODEX_RS_SSE_FIXTURE; use crate::flags::OPENAI_REQUEST_MAX_RETRIES; use crate::flags::OPENAI_STREAM_IDLE_TIMEOUT_MS; use crate::model_provider_info::ModelProviderInfo; +use crate::model_provider_info::WireApi; use crate::models::ResponseItem; use crate::util::backoff; -/// API request payload for a single model turn. -#[derive(Default, Debug, Clone)] -pub struct Prompt { - /// Conversation context input items. - pub input: Vec, - /// Optional previous response ID (when storage is enabled). - pub prev_id: Option, - /// Optional initial instructions (only sent on first turn). - pub instructions: Option, - /// Whether to store response on server side (disable_response_storage = !store). - pub store: bool, - - /// Additional tools sourced from external MCP servers. Note each key is - /// the "fully qualified" tool name (i.e., prefixed with the server name), - /// which should be reported to the model in place of Tool::name. - pub extra_tools: HashMap, -} - -#[derive(Debug)] -pub enum ResponseEvent { - OutputItemDone(ResponseItem), - Completed { response_id: String }, -} - -#[derive(Debug, Serialize)] -struct Payload<'a> { - model: &'a str, - #[serde(skip_serializing_if = "Option::is_none")] - instructions: Option<&'a String>, - // TODO(mbolin): ResponseItem::Other should not be serialized. Currently, - // we code defensively to avoid this case, but perhaps we should use a - // separate enum for serialization. - input: &'a Vec, - tools: &'a [serde_json::Value], - tool_choice: &'static str, - parallel_tool_calls: bool, - reasoning: Option, - #[serde(skip_serializing_if = "Option::is_none")] - previous_response_id: Option, - /// true when using the Responses API. - store: bool, - stream: bool, -} - -#[derive(Debug, Serialize)] -struct Reasoning { - effort: &'static str, - #[serde(skip_serializing_if = "Option::is_none")] - generate_summary: Option, -} - /// When serialized as JSON, this produces a valid "Tool" in the OpenAI /// Responses API. #[derive(Debug, Serialize)] @@ -152,7 +104,20 @@ impl ModelClient { } } - pub async fn stream(&mut self, prompt: &Prompt) -> Result { + /// Dispatches to either the Responses or Chat implementation depending on + /// the provider config. Public callers always invoke `stream()` – the + /// specialised helpers are private to avoid accidental misuse. + pub async fn stream(&self, prompt: &Prompt) -> Result { + match self.provider.wire_api { + WireApi::Responses => self.stream_responses(prompt).await, + WireApi::Chat => { + stream_chat_completions(prompt, &self.model, &self.client, &self.provider).await + } + } + } + + /// Implementation for the OpenAI *Responses* experimental API. + async fn stream_responses(&self, prompt: &Prompt) -> Result { if let Some(path) = &*CODEX_RS_SSE_FIXTURE { // short circuit for tests warn!(path, "Streaming from fixture"); @@ -200,10 +165,7 @@ impl ModelClient { loop { attempt += 1; - let api_key = self - .provider - .api_key() - .ok_or_else(|| crate::error::CodexErr::EnvVar("API_KEY"))?; + let api_key = self.provider.api_key()?; let res = self .client .post(&url) @@ -396,18 +358,6 @@ where } } -pub struct ResponseStream { - rx_event: mpsc::Receiver>, -} - -impl Stream for ResponseStream { - type Item = Result; - - fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - self.rx_event.poll_recv(cx) - } -} - /// used in tests to stream from a text SSE file async fn stream_from_fixture(path: impl AsRef) -> Result { let (tx_event, rx_event) = mpsc::channel::>(16); diff --git a/codex-rs/core/src/client_common.rs b/codex-rs/core/src/client_common.rs new file mode 100644 index 0000000000..514b6b60a8 --- /dev/null +++ b/codex-rs/core/src/client_common.rs @@ -0,0 +1,72 @@ +use crate::error::Result; +use crate::models::ResponseItem; +use futures::Stream; +use serde::Serialize; +use std::collections::HashMap; +use std::pin::Pin; +use std::task::Context; +use std::task::Poll; +use tokio::sync::mpsc; + +/// API request payload for a single model turn. +#[derive(Default, Debug, Clone)] +pub struct Prompt { + /// Conversation context input items. + pub input: Vec, + /// Optional previous response ID (when storage is enabled). + pub prev_id: Option, + /// Optional initial instructions (only sent on first turn). + pub instructions: Option, + /// Whether to store response on server side (disable_response_storage = !store). + pub store: bool, + + /// Additional tools sourced from external MCP servers. Note each key is + /// the "fully qualified" tool name (i.e., prefixed with the server name), + /// which should be reported to the model in place of Tool::name. + pub extra_tools: HashMap, +} + +#[derive(Debug)] +pub enum ResponseEvent { + OutputItemDone(ResponseItem), + Completed { response_id: String }, +} + +#[derive(Debug, Serialize)] +pub(crate) struct Reasoning { + pub(crate) effort: &'static str, + #[serde(skip_serializing_if = "Option::is_none")] + pub(crate) generate_summary: Option, +} + +#[derive(Debug, Serialize)] +pub(crate) struct Payload<'a> { + pub(crate) model: &'a str, + #[serde(skip_serializing_if = "Option::is_none")] + pub(crate) instructions: Option<&'a String>, + // TODO(mbolin): ResponseItem::Other should not be serialized. Currently, + // we code defensively to avoid this case, but perhaps we should use a + // separate enum for serialization. + pub(crate) input: &'a Vec, + pub(crate) tools: &'a [serde_json::Value], + pub(crate) tool_choice: &'static str, + pub(crate) parallel_tool_calls: bool, + pub(crate) reasoning: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub(crate) previous_response_id: Option, + /// true when using the Responses API. + pub(crate) store: bool, + pub(crate) stream: bool, +} + +pub(crate) struct ResponseStream { + pub(crate) rx_event: mpsc::Receiver>, +} + +impl Stream for ResponseStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.rx_event.poll_recv(cx) + } +} diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index 039e11ce9e..2957abb20b 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -29,8 +29,8 @@ use tracing::trace; use tracing::warn; use crate::client::ModelClient; -use crate::client::Prompt; -use crate::client::ResponseEvent; +use crate::client_common::Prompt; +use crate::client_common::ResponseEvent; use crate::config::Config; use crate::error::CodexErr; use crate::error::Result as CodexResult; diff --git a/codex-rs/core/src/error.rs b/codex-rs/core/src/error.rs index 0e438700cc..21431d2a69 100644 --- a/codex-rs/core/src/error.rs +++ b/codex-rs/core/src/error.rs @@ -98,7 +98,7 @@ pub enum CodexErr { TokioJoin(#[from] JoinError), #[error("missing environment variable {0}")] - EnvVar(&'static str), + EnvVar(String), } impl CodexErr { diff --git a/codex-rs/core/src/lib.rs b/codex-rs/core/src/lib.rs index 1c3a46dfd1..c81c91e75c 100644 --- a/codex-rs/core/src/lib.rs +++ b/codex-rs/core/src/lib.rs @@ -5,7 +5,9 @@ // the TUI or the tracing stack). #![deny(clippy::print_stdout, clippy::print_stderr)] +mod chat_completions; mod client; +mod client_common; pub mod codex; pub use codex::Codex; pub mod codex_wrapper; @@ -21,6 +23,7 @@ pub mod mcp_server_config; mod mcp_tool_call; mod model_provider_info; pub use model_provider_info::ModelProviderInfo; +pub use model_provider_info::WireApi; mod models; pub mod protocol; mod rollout; diff --git a/codex-rs/core/src/model_provider_info.rs b/codex-rs/core/src/model_provider_info.rs index e7069c0460..571640ea4e 100644 --- a/codex-rs/core/src/model_provider_info.rs +++ b/codex-rs/core/src/model_provider_info.rs @@ -9,6 +9,22 @@ use serde::Deserialize; use serde::Serialize; use std::collections::HashMap; +/// Wire protocol that the provider speaks. Most third-party services only +/// implement the classic OpenAI Chat Completions JSON schema, whereas OpenAI +/// itself (and a handful of others) additionally expose the more modern +/// *Responses* API. The two protocols use different request/response shapes +/// and *cannot* be auto-detected at runtime, therefore each provider entry +/// must declare which one it expects. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum WireApi { + /// The experimental “Responses” API exposed by OpenAI at `/v1/responses`. + #[default] + Responses, + /// Regular Chat Completions compatible with `/v1/chat/completions`. + Chat, +} + /// Serializable representation of a provider definition. #[derive(Debug, Clone, Deserialize, Serialize)] pub struct ModelProviderInfo { @@ -18,12 +34,19 @@ pub struct ModelProviderInfo { pub base_url: String, /// Environment variable that stores the user's API key for this provider. pub env_key: String, + + /// Which wire protocol this provider expects. Defaults to + /// `WireApi::Responses` to keep backward-compatibility with existing user + /// configs. + #[serde(default)] + pub wire_api: WireApi, } impl ModelProviderInfo { /// Returns the API key for this provider if present in the environment. - pub fn api_key(&self) -> Option { - std::env::var(&self.env_key).ok() + pub fn api_key(&self) -> crate::error::Result { + std::env::var(&self.env_key) + .map_err(|_| crate::error::CodexErr::EnvVar(self.env_key.clone())) } } @@ -38,6 +61,7 @@ pub fn built_in_model_providers() -> HashMap { name: "OpenAI".into(), base_url: "https://api.openai.com/v1".into(), env_key: "OPENAI_API_KEY".into(), + wire_api: WireApi::Responses, }, ), ( @@ -46,6 +70,7 @@ pub fn built_in_model_providers() -> HashMap { name: "OpenRouter".into(), base_url: "https://openrouter.ai/api/v1".into(), env_key: "OPENROUTER_API_KEY".into(), + wire_api: WireApi::Chat, }, ), ( @@ -54,6 +79,7 @@ pub fn built_in_model_providers() -> HashMap { name: "Gemini".into(), base_url: "https://generativelanguage.googleapis.com/v1beta/openai".into(), env_key: "GEMINI_API_KEY".into(), + wire_api: WireApi::Chat, }, ), ( @@ -62,6 +88,7 @@ pub fn built_in_model_providers() -> HashMap { name: "Ollama".into(), base_url: "http://localhost:11434/v1".into(), env_key: "OLLAMA_API_KEY".into(), + wire_api: WireApi::Chat, }, ), ( @@ -70,6 +97,7 @@ pub fn built_in_model_providers() -> HashMap { name: "Mistral".into(), base_url: "https://api.mistral.ai/v1".into(), env_key: "MISTRAL_API_KEY".into(), + wire_api: WireApi::Chat, }, ), ( @@ -78,6 +106,7 @@ pub fn built_in_model_providers() -> HashMap { name: "DeepSeek".into(), base_url: "https://api.deepseek.com".into(), env_key: "DEEPSEEK_API_KEY".into(), + wire_api: WireApi::Chat, }, ), ( @@ -86,6 +115,7 @@ pub fn built_in_model_providers() -> HashMap { name: "xAI".into(), base_url: "https://api.x.ai/v1".into(), env_key: "XAI_API_KEY".into(), + wire_api: WireApi::Chat, }, ), ( @@ -94,6 +124,7 @@ pub fn built_in_model_providers() -> HashMap { name: "Groq".into(), base_url: "https://api.groq.com/openai/v1".into(), env_key: "GROQ_API_KEY".into(), + wire_api: WireApi::Chat, }, ), ] diff --git a/codex-rs/core/tests/previous_response_id.rs b/codex-rs/core/tests/previous_response_id.rs index 50c1ba39ea..5bd60a9889 100644 --- a/codex-rs/core/tests/previous_response_id.rs +++ b/codex-rs/core/tests/previous_response_id.rs @@ -91,6 +91,7 @@ async fn keeps_previous_response_id_between_tasks() { // ModelClient will return an error if the environment variable for the // provider is not set. env_key: "PATH".into(), + wire_api: codex_core::WireApi::Responses, }; // Init session diff --git a/codex-rs/core/tests/stream_no_completed.rs b/codex-rs/core/tests/stream_no_completed.rs index 1af5fc4a56..996470f9ea 100644 --- a/codex-rs/core/tests/stream_no_completed.rs +++ b/codex-rs/core/tests/stream_no_completed.rs @@ -81,6 +81,7 @@ async fn retries_on_early_close() { // ModelClient will return an error if the environment variable for the // provider is not set. env_key: "PATH".into(), + wire_api: codex_core::WireApi::Responses, }; let ctrl_c = std::sync::Arc::new(tokio::sync::Notify::new());