diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 0c12880c94..64e51bb05e 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -1076,6 +1076,7 @@ dependencies = [ "thiserror 2.0.16", "time", "tokio", + "tokio-stream", "tokio-test", "tokio-util", "toml", diff --git a/codex-rs/core/Cargo.toml b/codex-rs/core/Cargo.toml index 4259e64fc4..af63d9058f 100644 --- a/codex-rs/core/Cargo.toml +++ b/codex-rs/core/Cargo.toml @@ -61,6 +61,7 @@ tokio = { workspace = true, features = [ "rt-multi-thread", "signal", ] } +tokio-stream = { workspace = true } tokio-util = { workspace = true, features = ["rt"] } toml = { workspace = true } toml_edit = { workspace = true } diff --git a/codex-rs/core/src/chat_completions.rs b/codex-rs/core/src/chat_completions.rs index d6f394fb86..2a7a01bbb2 100644 --- a/codex-rs/core/src/chat_completions.rs +++ b/codex-rs/core/src/chat_completions.rs @@ -31,6 +31,7 @@ use tracing::debug; use tracing::trace; /// Implementation for the classic Chat Completions API. +#[allow(dead_code)] pub(crate) async fn stream_chat_completions( prompt: &Prompt, model_family: &ModelFamily, @@ -361,6 +362,7 @@ pub(crate) async fn stream_chat_completions( /// 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. +#[allow(dead_code)] async fn process_chat_sse( stream: S, tx_event: mpsc::Sender>, @@ -660,6 +662,7 @@ async fn process_chat_sse( /// [`AggregateStreamExt::aggregate()`] keep receiving the original unmodified /// events. #[derive(Copy, Clone, Eq, PartialEq)] +#[allow(dead_code)] enum AggregateMode { AggregatedOnly, Streaming, @@ -885,6 +888,7 @@ impl AggregatedChatStream { } } + #[allow(dead_code)] pub(crate) fn streaming_mode(inner: S) -> Self { Self::new(inner, AggregateMode::Streaming) } diff --git a/codex-rs/core/src/client.rs b/codex-rs/core/src/client.rs index 3ea2ca79b5..ad4f2922fc 100644 --- a/codex-rs/core/src/client.rs +++ b/codex-rs/core/src/client.rs @@ -1,35 +1,20 @@ +use std::collections::HashMap; +use std::fs::File; use std::io::BufRead; -use std::path::Path; +use std::io::BufReader; +use std::io::Cursor; +use std::pin::Pin; +use std::sync::Arc; use std::sync::OnceLock; +use std::task::Context; +use std::task::Poll; use std::time::Duration; +use std::time::Instant; -use crate::AuthManager; use crate::auth::CodexAuth; -use crate::error::RetryLimitReachedError; -use crate::error::UnexpectedResponseError; -use bytes::Bytes; -use codex_app_server_protocol::AuthMode; -use codex_protocol::ConversationId; -use eventsource_stream::Eventsource; -use futures::prelude::*; -use regex_lite::Regex; -use reqwest::StatusCode; -use reqwest::header::HeaderMap; -use serde::Deserialize; -use serde::Serialize; -use serde_json::Value; -use tokio::sync::mpsc; -use tokio::time::timeout; -use tokio_util::io::ReaderStream; -use tracing::debug; -use tracing::trace; -use tracing::warn; - use crate::chat_completions::AggregateStreamExt; -use crate::chat_completions::stream_chat_completions; use crate::client_common::Prompt; use crate::client_common::ResponseEvent; -use crate::client_common::ResponseStream; use crate::client_common::ResponsesApiRequest; use crate::client_common::create_reasoning_param_for_request; use crate::client_common::create_text_param_for_request; @@ -37,53 +22,120 @@ use crate::config::Config; use crate::default_client::create_client; use crate::error::CodexErr; use crate::error::Result; +use crate::error::RetryLimitReachedError; +use crate::error::UnexpectedResponseError; use crate::error::UsageLimitReachedError; use crate::flags::CODEX_RS_SSE_FIXTURE; -use crate::model_family::ModelFamily; use crate::model_provider_info::ModelProviderInfo; use crate::model_provider_info::WireApi; -use crate::openai_model_info::get_model_info; -use crate::openai_tools::create_tools_json_for_responses_api; use crate::protocol::RateLimitSnapshot; use crate::protocol::RateLimitWindow; use crate::protocol::TokenUsage; use crate::token_data::PlanType; -use crate::util::backoff; +use bytes::Bytes; +use codex_app_server_protocol::AuthMode; use codex_otel::otel_event_manager::OtelEventManager; +use codex_protocol::ConversationId; use codex_protocol::config_types::ReasoningEffort as ReasoningEffortConfig; use codex_protocol::config_types::ReasoningSummary as ReasoningSummaryConfig; +use codex_protocol::models::ContentItem; +use codex_protocol::models::ReasoningItemContent; use codex_protocol::models::ResponseItem; -use std::sync::Arc; +use eventsource_stream::Eventsource; +use futures::Stream; +use futures::StreamExt; +use futures::TryStreamExt; +use regex_lite::Regex; +use reqwest::StatusCode; +use reqwest::header::HeaderMap; +use serde::Deserialize; +use serde::Serialize; +use serde_json::Value; +use serde_json::json; +use tokio::sync::mpsc; +use tokio::time::sleep; +use tokio::time::timeout; +use tokio_stream::wrappers::ReceiverStream; +use tokio_util::io::ReaderStream; +use tracing::debug; +use tracing::trace; +use tracing::warn; -#[derive(Debug, Deserialize)] -struct ErrorResponse { - error: Error, +use crate::AuthManager; +use crate::openai_tools::create_tools_json_for_chat_completions_api; +use crate::openai_tools::create_tools_json_for_responses_api; + +#[derive(Copy, Clone, Debug, Default, Eq, PartialEq)] +pub enum StreamMode { + #[default] + Aggregated, + Streaming, } -#[derive(Debug, Deserialize)] -struct Error { - r#type: Option, - code: Option, - message: Option, - - // Optional fields available on "usage_limit_reached" and "usage_not_included" errors - plan_type: Option, - resets_in_seconds: Option, +#[derive(Clone, Debug, Default)] +pub struct CallOpts { + pub stream_mode: Option, + pub conversation_id: Option, + pub provider_hint: Option, + pub effort: Option, + pub summary: Option, + pub output_schema: Option, + pub show_raw_reasoning: Option, } -#[derive(Debug, Clone)] -pub struct ModelClient { - config: Arc, - auth_manager: Option>, - otel_event_manager: OtelEventManager, - client: reqwest::Client, - provider: ModelProviderInfo, +#[derive(Clone, Debug)] +struct ResolvedCall { + stream_mode: StreamMode, conversation_id: ConversationId, + provider_hint: Option, effort: Option, summary: ReasoningSummaryConfig, + output_schema: Option, + show_raw_reasoning: bool, } -impl ModelClient { +#[derive(Clone, Debug)] +pub struct Client { + cfg: Arc, + provider: Arc, + auth: Option>, + http: reqwest::Client, + otel: OtelEventManager, + defaults: CallOpts, +} + +#[derive(Default)] +pub struct ClientBuilder { + config: Option>, + provider: Option, + auth: Option>, + otel: Option, + http: Option, + defaults: CallOpts, +} + +pub struct TurnStream { + inner: Pin> + Send + 'static>>, + dialect: WireDialect, +} + +pub struct TurnResult { + pub events: Vec, + pub error: Option, +} + +#[derive(Copy, Clone, Debug, Eq, PartialEq)] +pub enum WireDialect { + Responses, + Chat, +} + +impl Client { + pub fn builder() -> ClientBuilder { + ClientBuilder::default() + } + + #[allow(clippy::too_many_arguments)] pub fn new( config: Arc, auth_manager: Option>, @@ -93,136 +145,176 @@ impl ModelClient { summary: ReasoningSummaryConfig, conversation_id: ConversationId, ) -> Self { - let client = create_client(); + ClientBuilder::default() + .config(config) + .provider(provider) + .auth_manager(auth_manager) + .otel(otel_event_manager) + .conversation_id(conversation_id) + .reasoning_effort(effort) + .reasoning_summary(summary) + .build() + } - Self { - config, - auth_manager, - otel_event_manager, - client, - provider, - conversation_id, - effort, - summary, + pub async fn stream(&self, prompt: &Prompt, opts: impl Into) -> Result { + let call = self.resolve_call(opts.into())?; + let dialect = self.resolve_dialect(call.provider_hint); + + if matches!(dialect, WireDialect::Responses) + && let Some(path) = &*CODEX_RS_SSE_FIXTURE + { + return self.stream_from_fixture(path, &call).await; } + + let stream = match dialect { + WireDialect::Responses => self.stream_responses(prompt, &call).await?, + WireDialect::Chat => self.stream_chat(prompt, &call).await?, + }; + + Ok(stream) } - pub fn get_model_context_window(&self) -> Option { - self.config - .model_context_window - .or_else(|| get_model_info(&self.config.model_family).map(|info| info.context_window)) - } + pub async fn complete(&self, prompt: &Prompt, opts: impl Into) -> Result { + let mut stream = self.stream(prompt, opts).await?; + let mut events = Vec::new(); + let mut first_err: Option = None; - pub fn get_auto_compact_token_limit(&self) -> Option { - self.config.model_auto_compact_token_limit.or_else(|| { - get_model_info(&self.config.model_family).and_then(|info| info.auto_compact_token_limit) + while let Some(item) = stream.next().await { + match item { + Ok(event) => events.push(event), + Err(err) => { + if first_err.is_none() { + first_err = Some(err); + } + } + } + } + + Ok(TurnResult { + events, + error: first_err, }) } - /// 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 { + fn resolve_dialect(&self, hint: Option) -> WireDialect { + if let Some(h) = hint { + return h; + } match self.provider.wire_api { - WireApi::Responses => self.stream_responses(prompt).await, - WireApi::Chat => { - // Create the raw streaming connection first. - let response_stream = stream_chat_completions( - prompt, - &self.config.model_family, - &self.client, - &self.provider, - &self.otel_event_manager, - ) - .await?; - - // Wrap it with the aggregation adapter so callers see *only* - // the final assistant message per turn (matching the - // behaviour of the Responses API). - let mut aggregated = if self.config.show_raw_agent_reasoning { - crate::chat_completions::AggregatedChatStream::streaming_mode(response_stream) - } else { - response_stream.aggregate() - }; - - // Bridge the aggregated stream back into a standard - // `ResponseStream` by forwarding events through a channel. - let (tx, rx) = mpsc::channel::>(16); - - tokio::spawn(async move { - use futures::StreamExt; - while let Some(ev) = aggregated.next().await { - // Exit early if receiver hung up. - if tx.send(ev).await.is_err() { - break; - } - } - }); - - Ok(ResponseStream { rx_event: rx }) - } + WireApi::Responses => WireDialect::Responses, + WireApi::Chat => WireDialect::Chat, } } - /// 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"); - return stream_from_fixture( - path, - self.provider.clone(), - self.otel_event_manager.clone(), - ) - .await; - } + fn resolve_call(&self, opts: CallOpts) -> Result { + let stream_mode = opts + .stream_mode + .or(self.defaults.stream_mode) + .unwrap_or(StreamMode::Aggregated); - let auth_manager = self.auth_manager.clone(); + let conversation_id = opts + .conversation_id + .or(self.defaults.conversation_id) + .ok_or_else(|| CodexErr::Fatal("conversation_id must be provided".to_string()))?; - let full_instructions = prompt.get_full_instructions(&self.config.model_family); + let summary = opts.summary.or(self.defaults.summary).unwrap_or_default(); + + Ok(ResolvedCall { + stream_mode, + conversation_id, + provider_hint: opts.provider_hint.or(self.defaults.provider_hint), + effort: opts.effort.or(self.defaults.effort), + summary, + output_schema: opts.output_schema.or(self.defaults.output_schema.clone()), + show_raw_reasoning: opts + .show_raw_reasoning + .or(self.defaults.show_raw_reasoning) + .unwrap_or(false), + }) + } + + pub fn get_provider(&self) -> ModelProviderInfo { + (*self.provider).clone() + } + + pub fn get_otel_event_manager(&self) -> OtelEventManager { + self.otel.clone() + } + + pub fn get_model(&self) -> String { + self.cfg.model.clone() + } + + pub fn get_model_family(&self) -> crate::model_family::ModelFamily { + self.cfg.model_family.clone() + } + + pub fn get_model_context_window(&self) -> Option { + self.cfg.model_context_window.or_else(|| { + crate::openai_model_info::get_model_info(&self.cfg.model_family) + .map(|info| info.context_window) + }) + } + + pub fn get_auto_compact_token_limit(&self) -> Option { + self.cfg.model_auto_compact_token_limit.or_else(|| { + crate::openai_model_info::get_model_info(&self.cfg.model_family) + .and_then(|info| info.auto_compact_token_limit) + }) + } + + pub fn get_reasoning_effort(&self) -> Option { + self.defaults.effort + } + + pub fn get_reasoning_summary(&self) -> ReasoningSummaryConfig { + self.defaults + .summary + .unwrap_or(self.cfg.model_reasoning_summary) + } + + pub fn get_auth_manager(&self) -> Option> { + self.auth.clone() + } + + async fn stream_responses(&self, prompt: &Prompt, call: &ResolvedCall) -> Result { + let auth_manager = self.auth.clone(); + + let full_instructions = prompt.get_full_instructions(&self.cfg.model_family); let tools_json = create_tools_json_for_responses_api(&prompt.tools)?; - let reasoning = create_reasoning_param_for_request( - &self.config.model_family, - self.effort, - self.summary, - ); + let reasoning = + create_reasoning_param_for_request(&self.cfg.model_family, call.effort, call.summary); - let include: Vec = if reasoning.is_some() { + let include = if reasoning.is_some() { vec!["reasoning.encrypted_content".to_string()] } else { - vec![] + Vec::new() }; let input_with_instructions = prompt.get_formatted_input(); + let output_schema = call + .output_schema + .clone() + .or_else(|| prompt.output_schema.clone()); - let verbosity = match &self.config.model_family.family { - family if family == "gpt-5" => self.config.model_verbosity, + let verbosity = match &self.cfg.model_family.family { + family if family == "gpt-5" => self.cfg.model_verbosity, _ => { - if self.config.model_verbosity.is_some() { + if self.cfg.model_verbosity.is_some() { warn!( "model_verbosity is set but ignored for non-gpt-5 model family: {}", - self.config.model_family.family + self.cfg.model_family.family ); } - None } }; - // Only include `text.verbosity` for GPT-5 family models - let text = create_text_param_for_request(verbosity, &prompt.output_schema); - - // In general, we want to explicitly send `store: false` when using the Responses API, - // but in practice, the Azure Responses API rejects `store: false`: - // - // - If store = false and id is sent an error is thrown that ID is not found - // - If store = false and id is not sent an error is thrown that ID is required - // - // For Azure, we send `store: true` and preserve reasoning item IDs. + let text = create_text_param_for_request(verbosity, &output_schema); let azure_workaround = self.provider.is_azure_responses_endpoint(); let payload = ResponsesApiRequest { - model: &self.config.model, + model: &self.cfg.model, instructions: &full_instructions, input: &input_with_instructions, tools: &tools_json, @@ -232,7 +324,7 @@ impl ModelClient { store: azure_workaround, stream: true, include, - prompt_cache_key: Some(self.conversation_id.to_string()), + prompt_cache_key: Some(call.conversation_id.to_string()), text, }; @@ -244,55 +336,403 @@ impl ModelClient { let max_attempts = self.provider.request_max_retries(); for attempt in 0..=max_attempts { match self - .attempt_stream_responses(attempt, &payload_json, &auth_manager) + .attempt_stream_responses(attempt, &payload_json, call, auth_manager.as_ref()) .await { - Ok(stream) => { - return Ok(stream); + Ok(resp) => { + let headers = resp.headers().clone(); + let stream = resp.bytes_stream().map_err(CodexErr::Reqwest); + return Ok(self.build_turn_stream( + WireDialect::Responses, + stream, + Some(headers), + call, + )); } - Err(StreamAttemptError::Fatal(e)) => { - return Err(e); - } - Err(retryable_attempt_error) => { + Err(StreamAttemptError::Fatal(err)) => return Err(err), + Err(err) => { if attempt == max_attempts { - return Err(retryable_attempt_error.into_error()); + return Err(err.into_error()); } - - tokio::time::sleep(retryable_attempt_error.delay(attempt)).await; + let delay = err.delay(attempt); + sleep(delay).await; } } } - unreachable!("stream_responses_attempt should always return"); + unreachable!("stream_responses attempts should return within loop"); + } + + async fn stream_chat(&self, prompt: &Prompt, call: &ResolvedCall) -> Result { + if prompt.output_schema.is_some() || call.output_schema.is_some() { + return Err(CodexErr::UnsupportedOperation( + "output_schema is not supported for Chat Completions API".to_string(), + )); + } + + let mut messages = Vec::::new(); + let full_instructions = prompt.get_full_instructions(&self.cfg.model_family); + messages.push(json!({"role": "system", "content": full_instructions})); + + let input = prompt.get_formatted_input(); + + let mut reasoning_by_anchor_index: HashMap = HashMap::new(); + + let mut last_emitted_role: Option<&str> = None; + for item in &input { + match item { + ResponseItem::Message { role, .. } => last_emitted_role = Some(role.as_str()), + ResponseItem::FunctionCall { .. } | ResponseItem::LocalShellCall { .. } => { + last_emitted_role = Some("assistant") + } + ResponseItem::FunctionCallOutput { .. } => last_emitted_role = Some("tool"), + ResponseItem::Reasoning { .. } | ResponseItem::Other => {} + ResponseItem::CustomToolCall { .. } => {} + ResponseItem::CustomToolCallOutput { .. } => {} + ResponseItem::WebSearchCall { .. } => {} + } + } + + let mut last_user_index: Option = None; + for (idx, item) in input.iter().enumerate() { + if let ResponseItem::Message { role, .. } = item + && role == "user" + { + last_user_index = Some(idx); + } + } + + if !matches!(last_emitted_role, Some("user")) { + for (idx, item) in input.iter().enumerate() { + if let Some(u_idx) = last_user_index + && idx <= u_idx + { + continue; + } + + if let ResponseItem::Reasoning { + content: Some(items), + .. + } = item + { + let mut text = String::new(); + for c in items { + match c { + ReasoningItemContent::ReasoningText { text: t } + | ReasoningItemContent::Text { text: t } => text.push_str(t), + } + } + if text.trim().is_empty() { + continue; + } + + let mut attached = false; + if idx > 0 + && let ResponseItem::Message { role, .. } = &input[idx - 1] + && role == "assistant" + { + reasoning_by_anchor_index + .entry(idx - 1) + .and_modify(|v| v.push_str(&text)) + .or_insert(text.clone()); + attached = true; + } + + if !attached && idx + 1 < input.len() { + match &input[idx + 1] { + ResponseItem::FunctionCall { .. } + | ResponseItem::LocalShellCall { .. } => { + reasoning_by_anchor_index + .entry(idx + 1) + .and_modify(|v| v.push_str(&text)) + .or_insert(text.clone()); + } + ResponseItem::Message { role, .. } if role == "assistant" => { + reasoning_by_anchor_index + .entry(idx + 1) + .and_modify(|v| v.push_str(&text)) + .or_insert(text.clone()); + } + _ => {} + } + } + } + } + } + + let mut last_assistant_text: Option = None; + for (idx, item) in input.iter().enumerate() { + match item { + ResponseItem::Message { role, content, .. } => { + let mut text = String::new(); + for c in content { + match c { + ContentItem::InputText { text: t } + | ContentItem::OutputText { text: t } => text.push_str(t), + _ => {} + } + } + if role == "assistant" { + if let Some(prev) = &last_assistant_text + && prev == &text + { + continue; + } + last_assistant_text = Some(text.clone()); + } + + let mut msg = json!({"role": role, "content": text}); + if role == "assistant" + && let Some(reasoning) = reasoning_by_anchor_index.get(&idx) + && let Some(obj) = msg.as_object_mut() + { + obj.insert("reasoning".to_string(), json!(reasoning)); + } + messages.push(msg); + } + ResponseItem::FunctionCall { + name, + arguments, + call_id, + .. + } => { + let mut msg = json!({ + "role": "assistant", + "content": null, + "tool_calls": [{ + "id": call_id, + "type": "function", + "function": { + "name": name, + "arguments": arguments, + } + }] + }); + if let Some(reasoning) = reasoning_by_anchor_index.get(&idx) + && let Some(obj) = msg.as_object_mut() + { + obj.insert("reasoning".to_string(), json!(reasoning)); + } + messages.push(msg); + } + ResponseItem::LocalShellCall { + id, + call_id: _, + status, + action, + } => { + let mut msg = json!({ + "role": "assistant", + "content": null, + "tool_calls": [{ + "id": id.clone().unwrap_or_default(), + "type": "local_shell_call", + "status": status, + "action": action, + }] + }); + if let Some(reasoning) = reasoning_by_anchor_index.get(&idx) + && let Some(obj) = msg.as_object_mut() + { + obj.insert("reasoning".to_string(), json!(reasoning)); + } + messages.push(msg); + } + ResponseItem::FunctionCallOutput { call_id, output } => { + messages.push(json!({ + "role": "tool", + "tool_call_id": call_id, + "content": output.content, + })); + } + ResponseItem::CustomToolCall { + id, + call_id: _, + name, + input: tool_input, + status: _, + } => { + messages.push(json!({ + "role": "assistant", + "content": null, + "tool_calls": [{ + "id": id, + "type": "custom", + "custom": { + "name": name, + "input": tool_input, + } + }] + })); + } + ResponseItem::CustomToolCallOutput { call_id, output } => { + messages.push(json!({ + "role": "tool", + "tool_call_id": call_id, + "content": output, + })); + } + ResponseItem::Reasoning { .. } + | ResponseItem::WebSearchCall { .. } + | ResponseItem::Other => { + continue; + } + } + } + + let tools_json = create_tools_json_for_chat_completions_api(&prompt.tools)?; + let payload = json!({ + "model": self.cfg.model_family.slug, + "messages": messages, + "stream": true, + "tools": tools_json, + }); + + let mut attempt = 0; + let max_retries = self.provider.request_max_retries(); + loop { + attempt += 1; + + let req_builder = self + .provider + .create_request_builder(&self.http, &None) + .await?; + let res = self + .otel + .log_request(attempt, || { + req_builder + .header(reqwest::header::ACCEPT, "text/event-stream") + .json(&payload) + .send() + }) + .await; + + match res { + Ok(resp) if resp.status().is_success() => { + let headers = resp.headers().clone(); + let stream = resp.bytes_stream().map_err(CodexErr::Reqwest); + return Ok(self.build_turn_stream( + WireDialect::Chat, + stream, + Some(headers), + call, + )); + } + 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(UnexpectedResponseError { + status, + body, + request_id: None, + })); + } + + if attempt > max_retries { + return Err(CodexErr::RetryLimit(RetryLimitReachedError { + status, + request_id: None, + })); + } + + 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(|| crate::util::backoff(attempt)); + sleep(delay).await; + } + Err(err) => { + if attempt > max_retries { + return Err(err.into()); + } + let delay = crate::util::backoff(attempt); + sleep(delay).await; + } + } + } + } + + async fn stream_from_fixture(&self, path: &str, call: &ResolvedCall) -> Result { + let file = File::open(path).map_err(CodexErr::Io)?; + let reader = BufReader::new(file); + let mut body = String::new(); + + for line in reader.lines() { + body.push_str(&line.map_err(CodexErr::Io)?); + body.push_str("\n\n"); + } + + let cursor = Cursor::new(body); + let stream = ReaderStream::new(cursor).map_err(CodexErr::Io); + Ok(self.build_turn_stream(WireDialect::Responses, stream, None, call)) + } + + fn apply_stream_mode( + &self, + stream: TurnStream, + dialect: WireDialect, + call: &ResolvedCall, + ) -> TurnStream { + if !matches!(dialect, WireDialect::Chat) { + return stream; + } + + match call.stream_mode { + StreamMode::Aggregated if !call.show_raw_reasoning => { + let aggregated = stream.aggregate(); + TurnStream::new(aggregated, dialect) + } + _ => stream, + } + } + + fn build_turn_stream( + &self, + dialect: WireDialect, + stream: S, + headers: Option, + call: &ResolvedCall, + ) -> TurnStream + where + S: Stream> + Send + Unpin + 'static, + { + let (tx, rx) = mpsc::channel::>(1024); + let otel = self.otel.clone(); + let provider = self.provider.clone(); + let call_clone = call.clone(); + tokio::spawn(async move { + run_sse_loop(dialect, stream, headers, provider, otel, call_clone, tx).await; + }); + + let base = TurnStream::new(ReceiverStream::new(rx), dialect); + self.apply_stream_mode(base, dialect, call) } - /// Single attempt to start a streaming Responses API call. async fn attempt_stream_responses( &self, attempt: u64, payload_json: &Value, - auth_manager: &Option>, - ) -> std::result::Result { - // Always fetch the latest auth in case a prior attempt refreshed the token. - let auth = auth_manager.as_ref().and_then(|m| m.auth()); - - trace!( - "POST to {}: {:?}", - self.provider.get_full_url(&auth), - serde_json::to_string(payload_json) - ); + call: &ResolvedCall, + auth_manager: Option<&Arc>, + ) -> std::result::Result { + let auth = auth_manager.and_then(|manager| manager.auth()); let mut req_builder = self .provider - .create_request_builder(&self.client, &auth) + .create_request_builder(&self.http, &auth) .await .map_err(StreamAttemptError::Fatal)?; req_builder = req_builder .header("OpenAI-Beta", "responses=experimental") - // Send session_id for compatibility. - .header("conversation_id", self.conversation_id.to_string()) - .header("session_id", self.conversation_id.to_string()) + .header("conversation_id", call.conversation_id.to_string()) + .header("session_id", call.conversation_id.to_string()) .header(reqwest::header::ACCEPT, "text/event-stream") .json(payload_json); @@ -303,10 +743,7 @@ impl ModelClient { req_builder = req_builder.header("chatgpt-account-id", account_id); } - let res = self - .otel_event_manager - .log_request(attempt, || req_builder.send()) - .await; + let res = self.otel.log_request(attempt, || req_builder.send()).await; let mut request_id = None; if let Ok(resp) = &res { @@ -314,69 +751,32 @@ impl ModelClient { .headers() .get("cf-ray") .map(|v| v.to_str().unwrap_or_default().to_string()); - - trace!( - "Response status: {}, cf-ray: {:?}", - resp.status(), - request_id - ); } match res { - Ok(resp) if resp.status().is_success() => { - let (tx_event, rx_event) = mpsc::channel::>(1600); + Ok(resp) if resp.status().is_success() => Ok(resp), + Ok(resp) => { + let status = resp.status(); - if let Some(snapshot) = parse_rate_limit_snapshot(resp.headers()) - && tx_event - .send(Ok(ResponseEvent::RateLimits(snapshot))) - .await - .is_err() - { - debug!("receiver dropped rate limit snapshot event"); - } - - // spawn task to process SSE - let stream = resp.bytes_stream().map_err(CodexErr::Reqwest); - tokio::spawn(process_sse( - stream, - tx_event, - self.provider.stream_idle_timeout(), - self.otel_event_manager.clone(), - )); - - Ok(ResponseStream { rx_event }) - } - Ok(res) => { - let status = res.status(); - - // Pull out Retry‑After header if present. - let retry_after_secs = res + let retry_after = resp .headers() .get(reqwest::header::RETRY_AFTER) .and_then(|v| v.to_str().ok()) - .and_then(|s| s.parse::().ok()); - let retry_after = retry_after_secs.map(|s| Duration::from_millis(s * 1_000)); + .and_then(|s| s.parse::().ok()) + .map(|s| Duration::from_millis(s * 1_000)); if status == StatusCode::UNAUTHORIZED - && let Some(manager) = auth_manager.as_ref() + && let Some(manager) = auth_manager && manager.auth().is_some() { let _ = manager.refresh_token().await; } - // The OpenAI Responses endpoint returns structured JSON bodies even for 4xx/5xx - // errors. When we bubble early with only the HTTP status the caller sees an opaque - // "unexpected status 400 Bad Request" which makes debugging nearly impossible. - // Instead, read (and include) the response text so higher layers and users see the - // exact error message (e.g. "Unknown parameter: 'input[0].metadata'"). The body is - // small and this branch only runs on error paths so the extra allocation is - // negligible. if !(status == StatusCode::TOO_MANY_REQUESTS || status == StatusCode::UNAUTHORIZED || status.is_server_error()) { - // Surface the error body to callers. Use `unwrap_or_default` per Clippy. - let body = res.text().await.unwrap_or_default(); + let body = resp.text().await.unwrap_or_default(); return Err(StreamAttemptError::Fatal(CodexErr::UnexpectedStatus( UnexpectedResponseError { status, @@ -387,13 +787,10 @@ impl ModelClient { } if status == StatusCode::TOO_MANY_REQUESTS { - let rate_limit_snapshot = parse_rate_limit_snapshot(res.headers()); - let body = res.json::().await.ok(); + let rate_limit_snapshot = parse_rate_limit_snapshot(resp.headers()); + let body = resp.json::().await.ok(); if let Some(ErrorResponse { error }) = body { if error.r#type.as_deref() == Some("usage_limit_reached") { - // Prefer the plan_type provided in the error message if present - // because it's more up to date than the one encoded in the auth - // token. let plan_type = error .plan_type .or_else(|| auth.as_ref().and_then(CodexAuth::get_plan_type)); @@ -416,94 +813,162 @@ impl ModelClient { request_id, }) } - Err(e) => Err(StreamAttemptError::RetryableTransportError(e.into())), - } - } - - pub fn get_provider(&self) -> ModelProviderInfo { - self.provider.clone() - } - - pub fn get_otel_event_manager(&self) -> OtelEventManager { - self.otel_event_manager.clone() - } - - /// Returns the currently configured model slug. - pub fn get_model(&self) -> String { - self.config.model.clone() - } - - /// Returns the currently configured model family. - pub fn get_model_family(&self) -> ModelFamily { - self.config.model_family.clone() - } - - /// Returns the current reasoning effort setting. - pub fn get_reasoning_effort(&self) -> Option { - self.effort - } - - /// Returns the current reasoning summary setting. - pub fn get_reasoning_summary(&self) -> ReasoningSummaryConfig { - self.summary - } - - pub fn get_auth_manager(&self) -> Option> { - self.auth_manager.clone() - } -} - -enum StreamAttemptError { - RetryableHttpError { - status: StatusCode, - retry_after: Option, - request_id: Option, - }, - RetryableTransportError(CodexErr), - Fatal(CodexErr), -} - -impl StreamAttemptError { - /// attempt is 0-based. - fn delay(&self, attempt: u64) -> Duration { - // backoff() uses 1-based attempts. - let backoff_attempt = attempt + 1; - match self { - Self::RetryableHttpError { retry_after, .. } => { - retry_after.unwrap_or_else(|| backoff(backoff_attempt)) - } - Self::RetryableTransportError { .. } => backoff(backoff_attempt), - Self::Fatal(_) => { - // Should not be called on Fatal errors. - Duration::from_secs(0) - } - } - } - - fn into_error(self) -> CodexErr { - match self { - Self::RetryableHttpError { - status, request_id, .. - } => { - if status == StatusCode::INTERNAL_SERVER_ERROR { - CodexErr::InternalServerError - } else { - CodexErr::RetryLimit(RetryLimitReachedError { status, request_id }) - } - } - Self::RetryableTransportError(error) => error, - Self::Fatal(error) => error, + Err(err) => Err(StreamAttemptError::RetryableTransportError(err.into())), } } } -#[derive(Debug, Deserialize, Serialize)] -struct SseEvent { - #[serde(rename = "type")] - kind: String, - response: Option, - item: Option, - delta: Option, +impl ClientBuilder { + pub fn config(mut self, config: Arc) -> Self { + self.config = Some(config); + self + } + + pub fn provider(mut self, provider: ModelProviderInfo) -> Self { + self.provider = Some(provider); + self + } + + pub fn auth_manager(mut self, auth: Option>) -> Self { + self.auth = auth; + self + } + + pub fn otel(mut self, otel: OtelEventManager) -> Self { + self.otel = Some(otel); + self + } + + pub fn http_client(mut self, http: reqwest::Client) -> Self { + self.http = Some(http); + self + } + + pub fn defaults(mut self, defaults: CallOpts) -> Self { + self.defaults = defaults; + self + } + + pub fn conversation_id(mut self, conversation_id: ConversationId) -> Self { + self.defaults.conversation_id = Some(conversation_id); + self + } + + pub fn reasoning_effort(mut self, effort: Option) -> Self { + self.defaults.effort = effort; + self + } + + pub fn reasoning_summary(mut self, summary: ReasoningSummaryConfig) -> Self { + self.defaults.summary = Some(summary); + self + } + + pub fn stream_mode(mut self, mode: StreamMode) -> Self { + self.defaults.stream_mode = Some(mode); + self + } + + pub fn show_raw_reasoning(mut self, show: bool) -> Self { + self.defaults.show_raw_reasoning = Some(show); + self + } + + pub fn build(self) -> Client { + let cfg = match self.config { + Some(cfg) => cfg, + None => panic!("config must be provided before building Client"), + }; + let provider = match self.provider { + Some(provider) => provider, + None => panic!("provider must be provided before building Client"), + }; + let otel = match self.otel { + Some(otel) => otel, + None => panic!("otel event manager must be provided before building Client"), + }; + + let mut defaults = self.defaults; + if defaults.summary.is_none() { + defaults.summary = Some(cfg.model_reasoning_summary); + } + if defaults.show_raw_reasoning.is_none() { + defaults.show_raw_reasoning = Some(cfg.show_raw_agent_reasoning); + } + if defaults.stream_mode.is_none() { + defaults.stream_mode = Some(if cfg.show_raw_agent_reasoning { + StreamMode::Streaming + } else { + StreamMode::Aggregated + }); + } + + let http = self.http.unwrap_or_else(create_client); + + Client { + cfg, + provider: Arc::new(provider), + auth: self.auth, + http, + otel, + defaults, + } + } +} + +impl Stream for TurnStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.inner.as_mut().poll_next(cx) + } +} + +impl Unpin for TurnStream {} + +impl TurnStream { + fn new(stream: S, dialect: WireDialect) -> Self + where + S: Stream> + Send + 'static, + { + Self { + inner: Box::pin(stream), + dialect, + } + } + + pub fn dialect(&self) -> WireDialect { + self.dialect + } +} + +impl From<()> for CallOpts { + fn from(_: ()) -> Self { + CallOpts::default() + } +} + +impl From for CallOpts { + fn from(mode: StreamMode) -> Self { + CallOpts { + stream_mode: Some(mode), + ..CallOpts::default() + } + } +} + +#[derive(Debug, Deserialize)] +struct ErrorResponse { + error: Error, +} + +#[derive(Debug, Deserialize)] +struct Error { + r#type: Option, + code: Option, + message: Option, + plan_type: Option, + resets_in_seconds: Option, } #[derive(Debug, Deserialize)] @@ -549,11 +1014,613 @@ struct ResponseCompletedOutputTokensDetails { reasoning_tokens: u64, } +#[derive(Debug, Deserialize, Serialize)] +struct SseEvent { + #[serde(rename = "type")] + kind: String, + response: Option, + item: Option, + delta: Option, +} + +#[derive(Debug)] +enum StreamAttemptError { + RetryableHttpError { + status: StatusCode, + retry_after: Option, + request_id: Option, + }, + RetryableTransportError(CodexErr), + Fatal(CodexErr), +} + +impl StreamAttemptError { + fn delay(&self, attempt: u64) -> Duration { + let backoff_attempt = attempt + 1; + match self { + Self::RetryableHttpError { retry_after, .. } => { + retry_after.unwrap_or_else(|| crate::util::backoff(backoff_attempt)) + } + Self::RetryableTransportError { .. } => crate::util::backoff(backoff_attempt), + Self::Fatal(_) => Duration::from_secs(0), + } + } + + fn into_error(self) -> CodexErr { + match self { + Self::RetryableHttpError { + status, request_id, .. + } => { + if status == StatusCode::INTERNAL_SERVER_ERROR { + CodexErr::InternalServerError + } else { + CodexErr::RetryLimit(RetryLimitReachedError { status, request_id }) + } + } + Self::RetryableTransportError(error) => error, + Self::Fatal(error) => error, + } + } +} + +#[derive(Default)] +struct DecodeOutcome { + events: Vec, + errors: Vec, + completed: bool, +} + +struct FinalizeResult { + completed_emitted: bool, + error_emitted: bool, +} + +enum DecoderState { + Responses(ResponsesDecoderState), + Chat(ChatDecoderState), +} + +impl DecoderState { + fn new(dialect: WireDialect) -> Self { + match dialect { + WireDialect::Responses => DecoderState::Responses(ResponsesDecoderState::default()), + WireDialect::Chat => DecoderState::Chat(ChatDecoderState::default()), + } + } +} + +#[derive(Default)] +struct ResponsesDecoderState { + completed: Option, + error: Option, +} + +#[derive(Default)] +struct ChatDecoderState { + fn_call: FunctionCallState, + assistant_text: String, + reasoning_text: String, +} + +#[derive(Default)] +struct FunctionCallState { + name: Option, + arguments: String, + call_id: Option, + active: bool, +} + +async fn run_sse_loop( + dialect: WireDialect, + stream: S, + headers: Option, + provider: Arc, + otel: OtelEventManager, + _call: ResolvedCall, + mut tx: mpsc::Sender>, +) where + S: Stream> + Send + Unpin + 'static, +{ + let mut decoder = DecoderState::new(dialect); + let mut event_stream = stream.eventsource(); + let idle_timeout = provider.stream_idle_timeout(); + let mut saw_completed = false; + + if let Some(headers) = headers.as_ref() { + emit_ratelimit_snapshot(headers, &mut tx).await; + } + + loop { + let start = Instant::now(); + let next = timeout(idle_timeout, event_stream.next()).await; + let duration = start.elapsed(); + otel.log_sse_event(&next, duration); + + let sse = match next { + Ok(Some(Ok(ev))) => ev, + Ok(Some(Err(err))) => { + forward_err(CodexErr::Stream(err.to_string(), None), &mut tx).await; + break; + } + Ok(None) => break, + Err(_) => { + forward_err( + CodexErr::Stream("idle timeout waiting for SSE".into(), None), + &mut tx, + ) + .await; + break; + } + }; + + trace!("SSE event: {}", sse.data); + let outcome = decode_sse_line(&mut decoder, &sse.data); + + for err in outcome.errors { + forward_err(err, &mut tx).await; + if tx.is_closed() { + return; + } + } + + for event in outcome.events { + saw_completed |= matches!(event, ResponseEvent::Completed { .. }); + if forward_event(event, &mut tx).await { + return; + } + } + + if tx.is_closed() { + return; + } + } + + let finalize = finalize_decoder(&mut decoder, &mut tx, &otel).await; + saw_completed |= finalize.completed_emitted; + + if !saw_completed { + if matches!(dialect, WireDialect::Responses) && !finalize.error_emitted { + let err = CodexErr::Stream("stream closed before response.completed".into(), None); + otel.see_event_completed_failed(&err); + forward_err(err, &mut tx).await; + } + let _ = tx + .send(Ok(ResponseEvent::Completed { + response_id: String::new(), + token_usage: None, + })) + .await; + } +} + +fn decode_sse_line(state: &mut DecoderState, line: &str) -> DecodeOutcome { + match state { + DecoderState::Responses(inner) => decode_responses_line(inner, line), + DecoderState::Chat(inner) => decode_chat_line(inner, line), + } +} + +fn decode_responses_line(state: &mut ResponsesDecoderState, line: &str) -> DecodeOutcome { + let mut outcome = DecodeOutcome::default(); + let event: SseEvent = match serde_json::from_str(line) { + Ok(ev) => ev, + Err(err) => { + debug!("Failed to parse SSE event: {err}, data: {line}"); + return outcome; + } + }; + + match event.kind.as_str() { + "response.output_item.done" => { + if let Some(item_val) = event.item { + match serde_json::from_value::(item_val) { + Ok(item) => outcome.events.push(ResponseEvent::OutputItemDone(item)), + Err(err) => debug!("failed to parse ResponseItem from output_item.done: {err}"), + } + } + } + "response.output_text.delta" => { + if let Some(delta) = event.delta { + outcome.events.push(ResponseEvent::OutputTextDelta(delta)); + } + } + "response.reasoning_summary_text.delta" => { + if let Some(delta) = event.delta { + outcome + .events + .push(ResponseEvent::ReasoningSummaryDelta(delta)); + } + } + "response.reasoning_text.delta" => { + if let Some(delta) = event.delta { + outcome + .events + .push(ResponseEvent::ReasoningContentDelta(delta)); + } + } + "response.created" => { + if event.response.is_some() { + outcome.events.push(ResponseEvent::Created); + } + } + "response.failed" => { + if let Some(resp_val) = event.response { + let mut error = Some(CodexErr::Stream( + "response.failed event received".to_string(), + None, + )); + + if let Some(err_val) = resp_val.get("error") { + match serde_json::from_value::(err_val.clone()) { + Ok(parsed) => { + if is_context_window_error(&parsed) { + error = Some(CodexErr::ContextWindowExceeded); + } else { + let delay = try_parse_retry_after(&parsed); + let message = parsed.message.unwrap_or_default(); + error = Some(CodexErr::Stream(message, delay)); + } + } + Err(err) => { + let msg = format!("failed to parse ErrorResponse: {err}"); + debug!("{msg}"); + error = Some(CodexErr::Stream(msg, None)); + } + } + } + + state.error = error; + } + } + "response.completed" => { + if let Some(resp_val) = event.response { + match serde_json::from_value::(resp_val) { + Ok(completed) => state.completed = Some(completed), + Err(err) => { + let msg = format!("failed to parse ResponseCompleted: {err}"); + debug!("{msg}"); + state.error = Some(CodexErr::Stream(msg, None)); + } + } + } + } + "response.output_item.added" => { + if let Some(item) = event.item.as_ref() + && let Some(ty) = item.get("type").and_then(|v| v.as_str()) + && ty == "web_search_call" + { + let call_id = item + .get("id") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + outcome + .events + .push(ResponseEvent::WebSearchCallBegin { call_id }); + } + } + "response.reasoning_summary_part.added" => { + outcome + .events + .push(ResponseEvent::ReasoningSummaryPartAdded); + } + "response.reasoning_summary_text.done" + | "response.content_part.done" + | "response.function_call_arguments.delta" + | "response.custom_tool_call_input.delta" + | "response.custom_tool_call_input.done" + | "response.in_progress" + | "response.output_text.done" => {} + _ => {} + } + + outcome +} + +fn decode_chat_line(state: &mut ChatDecoderState, line: &str) -> DecodeOutcome { + let mut outcome = DecodeOutcome::default(); + if line.trim() == "[DONE]" { + if !state.assistant_text.is_empty() { + let item = ResponseItem::Message { + role: "assistant".to_string(), + content: vec![ContentItem::OutputText { + text: std::mem::take(&mut state.assistant_text), + }], + id: None, + }; + outcome.events.push(ResponseEvent::OutputItemDone(item)); + } + + if !state.reasoning_text.is_empty() { + let item = ResponseItem::Reasoning { + id: String::new(), + summary: Vec::new(), + content: Some(vec![ReasoningItemContent::ReasoningText { + text: std::mem::take(&mut state.reasoning_text), + }]), + encrypted_content: None, + }; + outcome.events.push(ResponseEvent::OutputItemDone(item)); + } + + outcome.events.push(ResponseEvent::Completed { + response_id: String::new(), + token_usage: None, + }); + outcome.completed = true; + state.fn_call = FunctionCallState::default(); + return outcome; + } + + let chunk: Value = match serde_json::from_str(line) { + Ok(value) => value, + Err(_) => return outcome, + }; + trace!("chat_completions received SSE chunk: {chunk:?}"); + + let Some(choice) = chunk.get("choices").and_then(|c| c.get(0)) else { + return outcome; + }; + + if let Some(content) = choice + .get("delta") + .and_then(|d| d.get("content")) + .and_then(|c| c.as_str()) + && !content.is_empty() + { + state.assistant_text.push_str(content); + outcome + .events + .push(ResponseEvent::OutputTextDelta(content.to_string())); + } + + if let Some(reasoning_val) = choice.get("delta").and_then(|d| d.get("reasoning")) { + let mut maybe_text = reasoning_val + .as_str() + .map(str::to_string) + .filter(|s| !s.is_empty()); + + if maybe_text.is_none() && reasoning_val.is_object() { + if let Some(s) = reasoning_val + .get("text") + .and_then(|t| t.as_str()) + .filter(|s| !s.is_empty()) + { + maybe_text = Some(s.to_string()); + } else if let Some(s) = reasoning_val + .get("content") + .and_then(|t| t.as_str()) + .filter(|s| !s.is_empty()) + { + maybe_text = Some(s.to_string()); + } + } + + if let Some(reasoning) = maybe_text { + state.reasoning_text.push_str(&reasoning); + outcome + .events + .push(ResponseEvent::ReasoningContentDelta(reasoning)); + } + } + + if let Some(message_reasoning) = choice.get("message").and_then(|m| m.get("reasoning")) { + if let Some(s) = message_reasoning.as_str() { + if !s.is_empty() { + state.reasoning_text.push_str(s); + outcome + .events + .push(ResponseEvent::ReasoningContentDelta(s.to_string())); + } + } else if let Some(obj) = message_reasoning.as_object() + && let Some(s) = obj + .get("text") + .and_then(|t| t.as_str()) + .or_else(|| obj.get("content").and_then(|t| t.as_str())) + .filter(|s| !s.is_empty()) + { + state.reasoning_text.push_str(s); + outcome + .events + .push(ResponseEvent::ReasoningContentDelta(s.to_string())); + } + } + + if let Some(tool_calls) = choice + .get("delta") + .and_then(|d| d.get("tool_calls")) + .and_then(|tc| tc.as_array()) + && let Some(tool_call) = tool_calls.first() + { + state.fn_call.active = true; + + if let Some(id) = tool_call.get("id").and_then(|v| v.as_str()) { + state.fn_call.call_id.get_or_insert_with(|| id.to_string()); + } + + if let Some(function) = tool_call.get("function") { + if let Some(name) = function.get("name").and_then(|n| n.as_str()) { + state.fn_call.name.get_or_insert_with(|| name.to_string()); + } + + if let Some(args_fragment) = function.get("arguments").and_then(|a| a.as_str()) { + state.fn_call.arguments.push_str(args_fragment); + } + } + } + + if let Some(finish_reason) = choice.get("finish_reason").and_then(|v| v.as_str()) { + match finish_reason { + "tool_calls" if state.fn_call.active => { + if !state.reasoning_text.is_empty() { + let item = ResponseItem::Reasoning { + id: String::new(), + summary: Vec::new(), + content: Some(vec![ReasoningItemContent::ReasoningText { + text: std::mem::take(&mut state.reasoning_text), + }]), + encrypted_content: None, + }; + outcome.events.push(ResponseEvent::OutputItemDone(item)); + } + + let item = ResponseItem::FunctionCall { + id: None, + name: state.fn_call.name.clone().unwrap_or_default(), + arguments: state.fn_call.arguments.clone(), + call_id: state.fn_call.call_id.clone().unwrap_or_default(), + }; + outcome.events.push(ResponseEvent::OutputItemDone(item)); + } + "stop" => { + if !state.reasoning_text.is_empty() { + let item = ResponseItem::Reasoning { + id: String::new(), + summary: Vec::new(), + content: Some(vec![ReasoningItemContent::ReasoningText { + text: std::mem::take(&mut state.reasoning_text), + }]), + encrypted_content: None, + }; + outcome.events.push(ResponseEvent::OutputItemDone(item)); + } + if !state.assistant_text.is_empty() { + let item = ResponseItem::Message { + role: "assistant".to_string(), + content: vec![ContentItem::OutputText { + text: std::mem::take(&mut state.assistant_text), + }], + id: None, + }; + outcome.events.push(ResponseEvent::OutputItemDone(item)); + } + } + _ => {} + } + + outcome.events.push(ResponseEvent::Completed { + response_id: String::new(), + token_usage: None, + }); + outcome.completed = true; + state.fn_call = FunctionCallState::default(); + } + + outcome +} + +async fn emit_ratelimit_snapshot( + headers: &HeaderMap, + tx: &mut mpsc::Sender>, +) { + if let Some(snapshot) = parse_rate_limit_snapshot(headers) { + let _ = tx.send(Ok(ResponseEvent::RateLimits(snapshot))).await; + } +} + +async fn finalize_decoder( + state: &mut DecoderState, + tx: &mut mpsc::Sender>, + otel: &OtelEventManager, +) -> FinalizeResult { + match state { + DecoderState::Responses(inner) => finalize_responses_state(inner, tx, otel).await, + DecoderState::Chat(inner) => finalize_chat_state(inner, tx).await, + } +} + +async fn finalize_responses_state( + state: &mut ResponsesDecoderState, + tx: &mut mpsc::Sender>, + otel: &OtelEventManager, +) -> FinalizeResult { + if let Some(completed) = state.completed.take() { + if let Some(usage) = &completed.usage { + otel.sse_event_completed( + usage.input_tokens, + usage.output_tokens, + usage.input_tokens_details.as_ref().map(|d| d.cached_tokens), + usage + .output_tokens_details + .as_ref() + .map(|d| d.reasoning_tokens), + usage.total_tokens, + ); + } + + let event = ResponseEvent::Completed { + response_id: completed.id, + token_usage: completed.usage.map(Into::into), + }; + let _ = tx.send(Ok(event)).await; + + return FinalizeResult { + completed_emitted: true, + error_emitted: false, + }; + } + + if let Some(error) = state.error.take() { + otel.see_event_completed_failed(&error); + let _ = tx.send(Err(error)).await; + return FinalizeResult { + completed_emitted: false, + error_emitted: true, + }; + } + + FinalizeResult { + completed_emitted: false, + error_emitted: false, + } +} + +async fn finalize_chat_state( + state: &mut ChatDecoderState, + tx: &mut mpsc::Sender>, +) -> FinalizeResult { + if !state.assistant_text.is_empty() { + let item = ResponseItem::Message { + role: "assistant".to_string(), + content: vec![ContentItem::OutputText { + text: std::mem::take(&mut state.assistant_text), + }], + id: None, + }; + let _ = tx.send(Ok(ResponseEvent::OutputItemDone(item))).await; + } + + if !state.reasoning_text.is_empty() { + let item = ResponseItem::Reasoning { + id: String::new(), + summary: Vec::new(), + content: Some(vec![ReasoningItemContent::ReasoningText { + text: std::mem::take(&mut state.reasoning_text), + }]), + encrypted_content: None, + }; + let _ = tx.send(Ok(ResponseEvent::OutputItemDone(item))).await; + } + + FinalizeResult { + completed_emitted: false, + error_emitted: false, + } +} + +async fn forward_event(event: ResponseEvent, tx: &mut mpsc::Sender>) -> bool { + tx.send(Ok(event)).await.is_err() +} + +async fn forward_err(err: CodexErr, tx: &mut mpsc::Sender>) { + let _ = tx.send(Err(err)).await; +} + fn attach_item_ids(payload_json: &mut Value, original_items: &[ResponseItem]) { let Some(input_value) = payload_json.get_mut("input") else { return; }; - let serde_json::Value::Array(items) = input_value else { + let Value::Array(items) = input_value else { return; }; @@ -633,266 +1700,6 @@ fn parse_header_str<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> { headers.get(name)?.to_str().ok() } -async fn process_sse( - stream: S, - tx_event: mpsc::Sender>, - idle_timeout: Duration, - otel_event_manager: OtelEventManager, -) where - S: Stream> + Unpin, -{ - let mut stream = stream.eventsource(); - - // If the stream stays completely silent for an extended period treat it as disconnected. - // The response id returned from the "complete" message. - let mut response_completed: Option = None; - let mut response_error: Option = None; - - loop { - let start = std::time::Instant::now(); - let response = timeout(idle_timeout, stream.next()).await; - let duration = start.elapsed(); - otel_event_manager.log_sse_event(&response, duration); - - let sse = match response { - Ok(Some(Ok(sse))) => sse, - Ok(Some(Err(e))) => { - debug!("SSE Error: {e:#}"); - let event = CodexErr::Stream(e.to_string(), None); - let _ = tx_event.send(Err(event)).await; - return; - } - Ok(None) => { - match response_completed { - Some(ResponseCompleted { - id: response_id, - usage, - }) => { - if let Some(token_usage) = &usage { - otel_event_manager.sse_event_completed( - token_usage.input_tokens, - token_usage.output_tokens, - token_usage - .input_tokens_details - .as_ref() - .map(|d| d.cached_tokens), - token_usage - .output_tokens_details - .as_ref() - .map(|d| d.reasoning_tokens), - token_usage.total_tokens, - ); - } - let event = ResponseEvent::Completed { - response_id, - token_usage: usage.map(Into::into), - }; - let _ = tx_event.send(Ok(event)).await; - } - None => { - let error = response_error.unwrap_or(CodexErr::Stream( - "stream closed before response.completed".into(), - None, - )); - otel_event_manager.see_event_completed_failed(&error); - - let _ = tx_event.send(Err(error)).await; - } - } - return; - } - Err(_) => { - let _ = tx_event - .send(Err(CodexErr::Stream( - "idle timeout waiting for SSE".into(), - None, - ))) - .await; - return; - } - }; - - let raw = sse.data.clone(); - trace!("SSE event: {}", raw); - - let event: SseEvent = match serde_json::from_str(&sse.data) { - Ok(event) => event, - Err(e) => { - debug!("Failed to parse SSE event: {e}, data: {}", &sse.data); - continue; - } - }; - - match event.kind.as_str() { - // Individual output item finalised. Forward immediately so the - // rest of the agent can stream assistant text/functions *live* - // instead of waiting for the final `response.completed` envelope. - // - // IMPORTANT: We used to ignore these events and forward the - // duplicated `output` array embedded in the `response.completed` - // payload. That produced two concrete issues: - // 1. No real‑time streaming – the user only saw output after the - // entire turn had finished, which broke the "typing" UX and - // made long‑running turns look stalled. - // 2. Duplicate `function_call_output` items – both the - // individual *and* the completed array were forwarded, which - // confused the backend and triggered 400 - // "previous_response_not_found" errors because the duplicated - // IDs did not match the incremental turn chain. - // - // The fix is to forward the incremental events *as they come* and - // drop the duplicated list inside `response.completed`. - "response.output_item.done" => { - let Some(item_val) = event.item else { continue }; - let Ok(item) = serde_json::from_value::(item_val) else { - debug!("failed to parse ResponseItem from output_item.done"); - continue; - }; - - let event = ResponseEvent::OutputItemDone(item); - if tx_event.send(Ok(event)).await.is_err() { - return; - } - } - "response.output_text.delta" => { - if let Some(delta) = event.delta { - let event = ResponseEvent::OutputTextDelta(delta); - if tx_event.send(Ok(event)).await.is_err() { - return; - } - } - } - "response.reasoning_summary_text.delta" => { - if let Some(delta) = event.delta { - let event = ResponseEvent::ReasoningSummaryDelta(delta); - if tx_event.send(Ok(event)).await.is_err() { - return; - } - } - } - "response.reasoning_text.delta" => { - if let Some(delta) = event.delta { - let event = ResponseEvent::ReasoningContentDelta(delta); - if tx_event.send(Ok(event)).await.is_err() { - return; - } - } - } - "response.created" => { - if event.response.is_some() { - let _ = tx_event.send(Ok(ResponseEvent::Created {})).await; - } - } - "response.failed" => { - if let Some(resp_val) = event.response { - response_error = Some(CodexErr::Stream( - "response.failed event received".to_string(), - None, - )); - - let error = resp_val.get("error"); - - if let Some(error) = error { - match serde_json::from_value::(error.clone()) { - Ok(error) => { - if is_context_window_error(&error) { - response_error = Some(CodexErr::ContextWindowExceeded); - } else { - let delay = try_parse_retry_after(&error); - let message = error.message.clone().unwrap_or_default(); - response_error = Some(CodexErr::Stream(message, delay)); - } - } - Err(e) => { - let error = format!("failed to parse ErrorResponse: {e}"); - debug!(error); - response_error = Some(CodexErr::Stream(error, None)) - } - } - } - } - } - // Final response completed – includes array of output items & id - "response.completed" => { - if let Some(resp_val) = event.response { - match serde_json::from_value::(resp_val) { - Ok(r) => { - response_completed = Some(r); - } - Err(e) => { - let error = format!("failed to parse ResponseCompleted: {e}"); - debug!(error); - response_error = Some(CodexErr::Stream(error, None)); - continue; - } - }; - }; - } - "response.content_part.done" - | "response.function_call_arguments.delta" - | "response.custom_tool_call_input.delta" - | "response.custom_tool_call_input.done" // also emitted as response.output_item.done - | "response.in_progress" - | "response.output_text.done" => {} - "response.output_item.added" => { - if let Some(item) = event.item.as_ref() { - // Detect web_search_call begin and forward a synthetic event upstream. - if let Some(ty) = item.get("type").and_then(|v| v.as_str()) - && ty == "web_search_call" - { - let call_id = item - .get("id") - .and_then(|v| v.as_str()) - .unwrap_or("") - .to_string(); - let ev = ResponseEvent::WebSearchCallBegin { call_id }; - if tx_event.send(Ok(ev)).await.is_err() { - return; - } - } - } - } - "response.reasoning_summary_part.added" => { - // Boundary between reasoning summary sections (e.g., titles). - let event = ResponseEvent::ReasoningSummaryPartAdded; - if tx_event.send(Ok(event)).await.is_err() { - return; - } - } - "response.reasoning_summary_text.done" => {} - _ => {} - } - } -} - -/// used in tests to stream from a text SSE file -async fn stream_from_fixture( - path: impl AsRef, - provider: ModelProviderInfo, - otel_event_manager: OtelEventManager, -) -> Result { - let (tx_event, rx_event) = mpsc::channel::>(1600); - let f = std::fs::File::open(path.as_ref())?; - let lines = std::io::BufReader::new(f).lines(); - - // insert \n\n after each line for proper SSE parsing - let mut content = String::new(); - for line in lines { - content.push_str(&line?); - content.push_str("\n\n"); - } - - let rdr = std::io::Cursor::new(content); - let stream = ReaderStream::new(rdr).map_err(CodexErr::Io); - tokio::spawn(process_sse( - stream, - tx_event, - provider.stream_idle_timeout(), - otel_event_manager, - )); - Ok(ResponseStream { rx_event }) -} - fn rate_limit_regex() -> &'static Regex { static RE: OnceLock = OnceLock::new(); @@ -905,7 +1712,6 @@ fn try_parse_retry_after(err: &Error) -> Option { return None; } - // parse the Please try again in 1.898s format using regex let re = rate_limit_regex(); if let Some(message) = &err.message && let Some(captures) = re.captures(message) @@ -931,487 +1737,4 @@ fn is_context_window_error(error: &Error) -> bool { error.code.as_deref() == Some("context_length_exceeded") } -#[cfg(test)] -mod tests { - use super::*; - use assert_matches::assert_matches; - use serde_json::json; - use tokio::sync::mpsc; - use tokio_test::io::Builder as IoBuilder; - use tokio_util::io::ReaderStream; - - // ──────────────────────────── - // Helpers - // ──────────────────────────── - - /// Runs the SSE parser on pre-chunked byte slices and returns every event - /// (including any final `Err` from a stream-closure check). - async fn collect_events( - chunks: &[&[u8]], - provider: ModelProviderInfo, - otel_event_manager: OtelEventManager, - ) -> Vec> { - let mut builder = IoBuilder::new(); - for chunk in chunks { - builder.read(chunk); - } - - let reader = builder.build(); - let stream = ReaderStream::new(reader).map_err(CodexErr::Io); - let (tx, mut rx) = mpsc::channel::>(16); - tokio::spawn(process_sse( - stream, - tx, - provider.stream_idle_timeout(), - otel_event_manager, - )); - - let mut events = Vec::new(); - while let Some(ev) = rx.recv().await { - events.push(ev); - } - events - } - - /// Builds an in-memory SSE stream from JSON fixtures and returns only the - /// successfully parsed events (panics on internal channel errors). - async fn run_sse( - events: Vec, - provider: ModelProviderInfo, - otel_event_manager: OtelEventManager, - ) -> Vec { - let mut body = String::new(); - for e in events { - let kind = e - .get("type") - .and_then(|v| v.as_str()) - .expect("fixture event missing type"); - if e.as_object().map(|o| o.len() == 1).unwrap_or(false) { - body.push_str(&format!("event: {kind}\n\n")); - } else { - body.push_str(&format!("event: {kind}\ndata: {e}\n\n")); - } - } - - let (tx, mut rx) = mpsc::channel::>(8); - let stream = ReaderStream::new(std::io::Cursor::new(body)).map_err(CodexErr::Io); - tokio::spawn(process_sse( - stream, - tx, - provider.stream_idle_timeout(), - otel_event_manager, - )); - - let mut out = Vec::new(); - while let Some(ev) = rx.recv().await { - out.push(ev.expect("channel closed")); - } - out - } - - fn otel_event_manager() -> OtelEventManager { - OtelEventManager::new( - ConversationId::new(), - "test", - "test", - None, - Some(AuthMode::ChatGPT), - false, - "test".to_string(), - ) - } - - // ──────────────────────────── - // Tests from `implement-test-for-responses-api-sse-parser` - // ──────────────────────────── - - #[tokio::test] - async fn parses_items_and_completed() { - let item1 = json!({ - "type": "response.output_item.done", - "item": { - "type": "message", - "role": "assistant", - "content": [{"type": "output_text", "text": "Hello"}] - } - }) - .to_string(); - - let item2 = json!({ - "type": "response.output_item.done", - "item": { - "type": "message", - "role": "assistant", - "content": [{"type": "output_text", "text": "World"}] - } - }) - .to_string(); - - let completed = json!({ - "type": "response.completed", - "response": { "id": "resp1" } - }) - .to_string(); - - let sse1 = format!("event: response.output_item.done\ndata: {item1}\n\n"); - let sse2 = format!("event: response.output_item.done\ndata: {item2}\n\n"); - let sse3 = format!("event: response.completed\ndata: {completed}\n\n"); - - let provider = ModelProviderInfo { - name: "test".to_string(), - base_url: Some("https://test.com".to_string()), - env_key: Some("TEST_API_KEY".to_string()), - env_key_instructions: None, - wire_api: WireApi::Responses, - query_params: None, - http_headers: None, - env_http_headers: None, - request_max_retries: Some(0), - stream_max_retries: Some(0), - stream_idle_timeout_ms: Some(1000), - requires_openai_auth: false, - }; - - let otel_event_manager = otel_event_manager(); - - let events = collect_events( - &[sse1.as_bytes(), sse2.as_bytes(), sse3.as_bytes()], - provider, - otel_event_manager, - ) - .await; - - assert_eq!(events.len(), 3); - - matches!( - &events[0], - Ok(ResponseEvent::OutputItemDone(ResponseItem::Message { role, .. })) - if role == "assistant" - ); - - matches!( - &events[1], - Ok(ResponseEvent::OutputItemDone(ResponseItem::Message { role, .. })) - if role == "assistant" - ); - - match &events[2] { - Ok(ResponseEvent::Completed { - response_id, - token_usage, - }) => { - assert_eq!(response_id, "resp1"); - assert!(token_usage.is_none()); - } - other => panic!("unexpected third event: {other:?}"), - } - } - - #[tokio::test] - async fn error_when_missing_completed() { - let item1 = json!({ - "type": "response.output_item.done", - "item": { - "type": "message", - "role": "assistant", - "content": [{"type": "output_text", "text": "Hello"}] - } - }) - .to_string(); - - let sse1 = format!("event: response.output_item.done\ndata: {item1}\n\n"); - let provider = ModelProviderInfo { - name: "test".to_string(), - base_url: Some("https://test.com".to_string()), - env_key: Some("TEST_API_KEY".to_string()), - env_key_instructions: None, - wire_api: WireApi::Responses, - query_params: None, - http_headers: None, - env_http_headers: None, - request_max_retries: Some(0), - stream_max_retries: Some(0), - stream_idle_timeout_ms: Some(1000), - requires_openai_auth: false, - }; - - let otel_event_manager = otel_event_manager(); - - let events = collect_events(&[sse1.as_bytes()], provider, otel_event_manager).await; - - assert_eq!(events.len(), 2); - - matches!(events[0], Ok(ResponseEvent::OutputItemDone(_))); - - match &events[1] { - Err(CodexErr::Stream(msg, _)) => { - assert_eq!(msg, "stream closed before response.completed") - } - other => panic!("unexpected second event: {other:?}"), - } - } - - #[tokio::test] - async fn error_when_error_event() { - let raw_error = r#"{"type":"response.failed","sequence_number":3,"response":{"id":"resp_689bcf18d7f08194bf3440ba62fe05d803fee0cdac429894","object":"response","created_at":1755041560,"status":"failed","background":false,"error":{"code":"rate_limit_exceeded","message":"Rate limit reached for gpt-5 in organization org-AAA on tokens per min (TPM): Limit 30000, Used 22999, Requested 12528. Please try again in 11.054s. Visit https://platform.openai.com/account/rate-limits to learn more."}, "usage":null,"user":null,"metadata":{}}}"#; - - let sse1 = format!("event: response.failed\ndata: {raw_error}\n\n"); - let provider = ModelProviderInfo { - name: "test".to_string(), - base_url: Some("https://test.com".to_string()), - env_key: Some("TEST_API_KEY".to_string()), - env_key_instructions: None, - wire_api: WireApi::Responses, - query_params: None, - http_headers: None, - env_http_headers: None, - request_max_retries: Some(0), - stream_max_retries: Some(0), - stream_idle_timeout_ms: Some(1000), - requires_openai_auth: false, - }; - - let otel_event_manager = otel_event_manager(); - - let events = collect_events(&[sse1.as_bytes()], provider, otel_event_manager).await; - - assert_eq!(events.len(), 1); - - match &events[0] { - Err(CodexErr::Stream(msg, delay)) => { - assert_eq!( - msg, - "Rate limit reached for gpt-5 in organization org-AAA on tokens per min (TPM): Limit 30000, Used 22999, Requested 12528. Please try again in 11.054s. Visit https://platform.openai.com/account/rate-limits to learn more." - ); - assert_eq!(*delay, Some(Duration::from_secs_f64(11.054))); - } - other => panic!("unexpected second event: {other:?}"), - } - } - - #[tokio::test] - async fn context_window_error_is_fatal() { - let raw_error = r#"{"type":"response.failed","sequence_number":3,"response":{"id":"resp_5c66275b97b9baef1ed95550adb3b7ec13b17aafd1d2f11b","object":"response","created_at":1759510079,"status":"failed","background":false,"error":{"code":"context_length_exceeded","message":"Your input exceeds the context window of this model. Please adjust your input and try again."},"usage":null,"user":null,"metadata":{}}}"#; - - let sse1 = format!("event: response.failed\ndata: {raw_error}\n\n"); - let provider = ModelProviderInfo { - name: "test".to_string(), - base_url: Some("https://test.com".to_string()), - env_key: Some("TEST_API_KEY".to_string()), - env_key_instructions: None, - wire_api: WireApi::Responses, - query_params: None, - http_headers: None, - env_http_headers: None, - request_max_retries: Some(0), - stream_max_retries: Some(0), - stream_idle_timeout_ms: Some(1000), - requires_openai_auth: false, - }; - - let otel_event_manager = otel_event_manager(); - - let events = collect_events(&[sse1.as_bytes()], provider, otel_event_manager).await; - - assert_eq!(events.len(), 1); - - match &events[0] { - Err(err @ CodexErr::ContextWindowExceeded) => { - assert_eq!(err.to_string(), CodexErr::ContextWindowExceeded.to_string()); - } - other => panic!("unexpected context window event: {other:?}"), - } - } - - #[tokio::test] - async fn context_window_error_with_newline_is_fatal() { - let raw_error = r#"{"type":"response.failed","sequence_number":4,"response":{"id":"resp_fatal_newline","object":"response","created_at":1759510080,"status":"failed","background":false,"error":{"code":"context_length_exceeded","message":"Your input exceeds the context window of this model. Please adjust your input and try\nagain."},"usage":null,"user":null,"metadata":{}}}"#; - - let sse1 = format!("event: response.failed\ndata: {raw_error}\n\n"); - let provider = ModelProviderInfo { - name: "test".to_string(), - base_url: Some("https://test.com".to_string()), - env_key: Some("TEST_API_KEY".to_string()), - env_key_instructions: None, - wire_api: WireApi::Responses, - query_params: None, - http_headers: None, - env_http_headers: None, - request_max_retries: Some(0), - stream_max_retries: Some(0), - stream_idle_timeout_ms: Some(1000), - requires_openai_auth: false, - }; - - let otel_event_manager = otel_event_manager(); - - let events = collect_events(&[sse1.as_bytes()], provider, otel_event_manager).await; - - assert_eq!(events.len(), 1); - - match &events[0] { - Err(err @ CodexErr::ContextWindowExceeded) => { - assert_eq!(err.to_string(), CodexErr::ContextWindowExceeded.to_string()); - } - other => panic!("unexpected context window event: {other:?}"), - } - } - - // ──────────────────────────── - // Table-driven test from `main` - // ──────────────────────────── - - /// Verifies that the adapter produces the right `ResponseEvent` for a - /// variety of incoming `type` values. - #[tokio::test] - async fn table_driven_event_kinds() { - struct TestCase { - name: &'static str, - event: serde_json::Value, - expect_first: fn(&ResponseEvent) -> bool, - expected_len: usize, - } - - fn is_created(ev: &ResponseEvent) -> bool { - matches!(ev, ResponseEvent::Created) - } - fn is_output(ev: &ResponseEvent) -> bool { - matches!(ev, ResponseEvent::OutputItemDone(_)) - } - fn is_completed(ev: &ResponseEvent) -> bool { - matches!(ev, ResponseEvent::Completed { .. }) - } - - let completed = json!({ - "type": "response.completed", - "response": { - "id": "c", - "usage": { - "input_tokens": 0, - "input_tokens_details": null, - "output_tokens": 0, - "output_tokens_details": null, - "total_tokens": 0 - }, - "output": [] - } - }); - - let cases = vec![ - TestCase { - name: "created", - event: json!({"type": "response.created", "response": {}}), - expect_first: is_created, - expected_len: 2, - }, - TestCase { - name: "output_item.done", - event: json!({ - "type": "response.output_item.done", - "item": { - "type": "message", - "role": "assistant", - "content": [ - {"type": "output_text", "text": "hi"} - ] - } - }), - expect_first: is_output, - expected_len: 2, - }, - TestCase { - name: "unknown", - event: json!({"type": "response.new_tool_event"}), - expect_first: is_completed, - expected_len: 1, - }, - ]; - - for case in cases { - let mut evs = vec![case.event]; - evs.push(completed.clone()); - - let provider = ModelProviderInfo { - name: "test".to_string(), - base_url: Some("https://test.com".to_string()), - env_key: Some("TEST_API_KEY".to_string()), - env_key_instructions: None, - wire_api: WireApi::Responses, - query_params: None, - http_headers: None, - env_http_headers: None, - request_max_retries: Some(0), - stream_max_retries: Some(0), - stream_idle_timeout_ms: Some(1000), - requires_openai_auth: false, - }; - - let otel_event_manager = otel_event_manager(); - - let out = run_sse(evs, provider, otel_event_manager).await; - assert_eq!(out.len(), case.expected_len, "case {}", case.name); - assert!( - (case.expect_first)(&out[0]), - "first event mismatch in case {}", - case.name - ); - } - } - - #[test] - fn test_try_parse_retry_after() { - let err = Error { - r#type: None, - message: Some("Rate limit reached for gpt-5 in organization org- on tokens per min (TPM): Limit 1, Used 1, Requested 19304. Please try again in 28ms. Visit https://platform.openai.com/account/rate-limits to learn more.".to_string()), - code: Some("rate_limit_exceeded".to_string()), - plan_type: None, - resets_in_seconds: None - }; - - let delay = try_parse_retry_after(&err); - assert_eq!(delay, Some(Duration::from_millis(28))); - } - - #[test] - fn test_try_parse_retry_after_no_delay() { - let err = Error { - r#type: None, - message: Some("Rate limit reached for gpt-5 in organization on tokens per min (TPM): Limit 30000, Used 6899, Requested 24050. Please try again in 1.898s. Visit https://platform.openai.com/account/rate-limits to learn more.".to_string()), - code: Some("rate_limit_exceeded".to_string()), - plan_type: None, - resets_in_seconds: None - }; - let delay = try_parse_retry_after(&err); - assert_eq!(delay, Some(Duration::from_secs_f64(1.898))); - } - - #[test] - fn error_response_deserializes_old_schema_known_plan_type_and_serializes_back() { - use crate::token_data::KnownPlan; - use crate::token_data::PlanType; - - let json = r#"{"error":{"type":"usage_limit_reached","plan_type":"pro","resets_in_seconds":3600}}"#; - let resp: ErrorResponse = - serde_json::from_str(json).expect("should deserialize old schema"); - - assert_matches!(resp.error.plan_type, Some(PlanType::Known(KnownPlan::Pro))); - - let plan_json = serde_json::to_string(&resp.error.plan_type).expect("serialize plan_type"); - assert_eq!(plan_json, "\"pro\""); - } - - #[test] - fn error_response_deserializes_old_schema_unknown_plan_type_and_serializes_back() { - use crate::token_data::PlanType; - - let json = - r#"{"error":{"type":"usage_limit_reached","plan_type":"vip","resets_in_seconds":60}}"#; - let resp: ErrorResponse = - serde_json::from_str(json).expect("should deserialize old schema"); - - assert_matches!(resp.error.plan_type, Some(PlanType::Unknown(ref s)) if s == "vip"); - - let plan_json = serde_json::to_string(&resp.error.plan_type).expect("serialize plan_type"); - assert_eq!(plan_json, "\"vip\""); - } -} +pub type ModelClient = Client; diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index fe352f0103..1c0113e17e 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -2109,7 +2109,7 @@ async fn try_run_turn( summary: turn_context.client.get_reasoning_summary(), }); sess.persist_rollout_items(&[rollout_item]).await; - let mut stream = turn_context.client.clone().stream(&prompt).await?; + let mut stream = turn_context.client.clone().stream(&prompt, ()).await?; let tool_runtime = ToolCallRuntime::new( Arc::clone(&router), diff --git a/codex-rs/core/src/codex/compact.rs b/codex-rs/core/src/codex/compact.rs index d43e3abcbb..fc0f002d4a 100644 --- a/codex-rs/core/src/codex/compact.rs +++ b/codex-rs/core/src/codex/compact.rs @@ -258,7 +258,7 @@ async fn drain_to_completed( sub_id: &str, prompt: &Prompt, ) -> CodexResult<()> { - let mut stream = turn_context.client.clone().stream(prompt).await?; + let mut stream = turn_context.client.clone().stream(prompt, ()).await?; loop { let maybe_event = stream.next().await; let Some(event) = maybe_event else { diff --git a/codex-rs/core/src/lib.rs b/codex-rs/core/src/lib.rs index 201d8feb4d..7817de6e38 100644 --- a/codex-rs/core/src/lib.rs +++ b/codex-rs/core/src/lib.rs @@ -93,7 +93,13 @@ pub use codex_protocol::protocol; // as those in the protocol crate when constructing protocol messages. pub use codex_protocol::config_types as protocol_config_types; +pub use client::CallOpts; +pub use client::Client; pub use client::ModelClient; +pub use client::StreamMode; +pub use client::TurnResult; +pub use client::TurnStream; +pub use client::WireDialect; pub use client_common::Prompt; pub use client_common::REVIEW_PROMPT; pub use client_common::ResponseEvent; diff --git a/codex-rs/core/tests/chat_completions_payload.rs b/codex-rs/core/tests/chat_completions_payload.rs index 9e10b37822..b0ed9bd956 100644 --- a/codex-rs/core/tests/chat_completions_payload.rs +++ b/codex-rs/core/tests/chat_completions_payload.rs @@ -97,7 +97,7 @@ async fn run_request(input: Vec) -> Value { let mut prompt = Prompt::default(); prompt.input = input; - let mut stream = match client.stream(&prompt).await { + let mut stream = match client.stream(&prompt, ()).await { Ok(s) => s, Err(e) => panic!("stream chat failed: {e}"), }; diff --git a/codex-rs/core/tests/chat_completions_sse.rs b/codex-rs/core/tests/chat_completions_sse.rs index 1aab6ac38c..b595a529b7 100644 --- a/codex-rs/core/tests/chat_completions_sse.rs +++ b/codex-rs/core/tests/chat_completions_sse.rs @@ -102,7 +102,7 @@ async fn run_stream_with_bytes(sse_body: &[u8]) -> Vec { }], }]; - let mut stream = match client.stream(&prompt).await { + let mut stream = match client.stream(&prompt, ()).await { Ok(s) => s, Err(e) => panic!("stream chat failed: {e}"), }; @@ -115,6 +115,9 @@ async fn run_stream_with_bytes(sse_body: &[u8]) -> Vec { } } events + .into_iter() + .filter(|ev| !matches!(ev, ResponseEvent::RateLimits(_))) + .collect() } fn assert_message(item: &ResponseItem, expected: &str) { diff --git a/codex-rs/core/tests/suite/client.rs b/codex-rs/core/tests/suite/client.rs index eb14dabb81..29fd79a63e 100644 --- a/codex-rs/core/tests/suite/client.rs +++ b/codex-rs/core/tests/suite/client.rs @@ -724,7 +724,7 @@ async fn azure_responses_request_includes_store_and_reasoning_ids() { }); let mut stream = client - .stream(&prompt) + .stream(&prompt, ()) .await .expect("responses stream to start");