diff --git a/codex-rs/codex-api/src/common.rs b/codex-rs/codex-api/src/common.rs index 3bb59a46a3..0e49606c63 100644 --- a/codex-rs/codex-api/src/common.rs +++ b/codex-rs/codex-api/src/common.rs @@ -20,6 +20,25 @@ use tokio::sync::mpsc; pub const WS_REQUEST_HEADER_TRACEPARENT_CLIENT_METADATA_KEY: &str = "ws_request_header_traceparent"; pub const WS_REQUEST_HEADER_TRACESTATE_CLIENT_METADATA_KEY: &str = "ws_request_header_tracestate"; +pub(crate) fn insert_max_output_tokens( + body: &mut Value, + max_output_tokens: Option, +) -> Result<(), ApiError> { + let Some(max_output_tokens) = max_output_tokens else { + return Ok(()); + }; + let Some(body) = body.as_object_mut() else { + return Err(ApiError::Stream( + "failed to add max_output_tokens to responses request".to_string(), + )); + }; + body.insert( + "max_output_tokens".to_string(), + Value::from(max_output_tokens), + ); + Ok(()) +} + /// Canonical input payload for the compaction endpoint. #[derive(Debug, Clone, Serialize)] pub struct CompactionInput<'a> { @@ -180,8 +199,6 @@ pub struct ResponsesApiRequest { pub stream: bool, pub include: Vec, #[serde(skip_serializing_if = "Option::is_none")] - pub max_output_tokens: Option, - #[serde(skip_serializing_if = "Option::is_none")] pub service_tier: Option, #[serde(skip_serializing_if = "Option::is_none")] pub prompt_cache_key: Option, @@ -205,7 +222,6 @@ impl From<&ResponsesApiRequest> for ResponseCreateWsRequest { store: request.store, stream: request.stream, include: request.include.clone(), - max_output_tokens: request.max_output_tokens, service_tier: request.service_tier.clone(), prompt_cache_key: request.prompt_cache_key.clone(), text: request.text.clone(), @@ -231,8 +247,6 @@ pub struct ResponseCreateWsRequest { pub stream: bool, pub include: Vec, #[serde(skip_serializing_if = "Option::is_none")] - pub max_output_tokens: Option, - #[serde(skip_serializing_if = "Option::is_none")] pub service_tier: Option, #[serde(skip_serializing_if = "Option::is_none")] pub prompt_cache_key: Option, diff --git a/codex-rs/codex-api/src/endpoint/responses.rs b/codex-rs/codex-api/src/endpoint/responses.rs index cc1be2846a..96335a33a2 100644 --- a/codex-rs/codex-api/src/endpoint/responses.rs +++ b/codex-rs/codex-api/src/endpoint/responses.rs @@ -1,6 +1,7 @@ use crate::auth::SharedAuthProvider; use crate::common::ResponseStream; use crate::common::ResponsesApiRequest; +use crate::common::insert_max_output_tokens; use crate::endpoint::session::EndpointSession; use crate::error::ApiError; use crate::provider::Provider; @@ -71,6 +72,18 @@ impl ResponsesClient { &self, request: ResponsesApiRequest, options: ResponsesOptions, + ) -> Result { + self.stream_request_with_max_output_tokens( + request, options, /*max_output_tokens*/ None, + ) + .await + } + + pub async fn stream_request_with_max_output_tokens( + &self, + request: ResponsesApiRequest, + options: ResponsesOptions, + max_output_tokens: Option, ) -> Result { let ResponsesOptions { session_id, @@ -83,6 +96,7 @@ impl ResponsesClient { let mut body = serde_json::to_value(&request) .map_err(|e| ApiError::Stream(format!("failed to encode responses request: {e}")))?; + insert_max_output_tokens(&mut body, max_output_tokens)?; if request.store && self.session.provider().is_azure_responses_endpoint() { attach_item_ids(&mut body, &request.input); } diff --git a/codex-rs/codex-api/src/endpoint/responses_websocket.rs b/codex-rs/codex-api/src/endpoint/responses_websocket.rs index f0a0019817..7859a02061 100644 --- a/codex-rs/codex-api/src/endpoint/responses_websocket.rs +++ b/codex-rs/codex-api/src/endpoint/responses_websocket.rs @@ -3,6 +3,7 @@ use crate::common::ResponseEvent; use crate::common::ResponseProcessedWsRequest; use crate::common::ResponseStream; use crate::common::ResponsesWsRequest; +use crate::common::insert_max_output_tokens; use crate::error::ApiError; use crate::provider::Provider; use crate::rate_limits::parse_rate_limit_event; @@ -250,6 +251,20 @@ impl ResponsesWebsocketConnection { &self, request: ResponsesWsRequest, connection_reused: bool, + ) -> Result { + self.stream_request_with_max_output_tokens( + request, + connection_reused, + /*max_output_tokens*/ None, + ) + .await + } + + pub async fn stream_request_with_max_output_tokens( + &self, + request: ResponsesWsRequest, + connection_reused: bool, + max_output_tokens: Option, ) -> Result { let (tx_event, rx_event) = mpsc::channel::>(1600); @@ -259,9 +274,10 @@ impl ResponsesWebsocketConnection { let models_etag = self.models_etag.clone(); let server_model = self.server_model.clone(); let telemetry = self.telemetry.clone(); - let request_body = serde_json::to_value(&request).map_err(|err| { + let mut request_body = serde_json::to_value(&request).map_err(|err| { ApiError::Stream(format!("failed to encode websocket request: {err}")) })?; + insert_max_output_tokens(&mut request_body, max_output_tokens)?; let current_span = Span::current(); tokio::spawn( diff --git a/codex-rs/codex-api/tests/clients.rs b/codex-rs/codex-api/tests/clients.rs index 2fe4c784e5..1a7e863b4c 100644 --- a/codex-rs/codex-api/tests/clients.rs +++ b/codex-rs/codex-api/tests/clients.rs @@ -331,7 +331,6 @@ async fn streaming_client_retries_on_transport_error() -> Result<()> { store: false, stream: true, include: Vec::new(), - max_output_tokens: None, service_tier: None, prompt_cache_key: None, text: None, @@ -433,7 +432,6 @@ async fn azure_default_store_attaches_ids_and_headers() -> Result<()> { store: true, stream: true, include: Vec::new(), - max_output_tokens: None, service_tier: None, prompt_cache_key: None, text: None, @@ -443,7 +441,7 @@ async fn azure_default_store_attaches_ids_and_headers() -> Result<()> { let mut extra_headers = HeaderMap::new(); extra_headers.insert("x-test-header", HeaderValue::from_static("present")); let _stream = client - .stream_request( + .stream_request_with_max_output_tokens( request, ResponsesOptions { session_id: Some("sess_123".into()), @@ -453,6 +451,7 @@ async fn azure_default_store_attaches_ids_and_headers() -> Result<()> { compression: Compression::None, turn_state: None, }, + /*max_output_tokens*/ Some(32), ) .await?; @@ -496,6 +495,13 @@ async fn azure_default_store_attaches_ids_and_headers() -> Result<()> { .and_then(|item| item.get("id")) .and_then(|id| id.as_str()); assert_eq!(input_id, Some("msg_1")); + let max_output_tokens = req + .body + .as_ref() + .and_then(RequestBody::json) + .and_then(|body| body.get("max_output_tokens")) + .and_then(serde_json::Value::as_u64); + assert_eq!(max_output_tokens, Some(32)); Ok(()) } diff --git a/codex-rs/core/src/client.rs b/codex-rs/core/src/client.rs index ebed7ea288..bceb250730 100644 --- a/codex-rs/core/src/client.rs +++ b/codex-rs/core/src/client.rs @@ -798,7 +798,6 @@ impl ModelClient { store: provider.is_azure_responses_endpoint(), stream: true, include, - max_output_tokens: prompt.max_output_tokens, service_tier, prompt_cache_key, text, @@ -1317,7 +1316,9 @@ impl ModelClientSession { client_setup.api_auth, ) .with_telemetry(Some(request_telemetry), Some(sse_telemetry)); - let stream_result = client.stream_request(request, options).await; + let stream_result = client + .stream_request_with_max_output_tokens(request, options, prompt.max_output_tokens) + .await; match stream_result { Ok(stream) => { @@ -1491,7 +1492,11 @@ impl ModelClientSession { )) })?; let stream_result = websocket_connection - .stream_request(ws_request, self.websocket_session.connection_reused()) + .stream_request_with_max_output_tokens( + ws_request, + self.websocket_session.connection_reused(), + prompt.max_output_tokens, + ) .await .map_err(|err| { let response_debug_context = diff --git a/codex-rs/core/src/client_common_tests.rs b/codex-rs/core/src/client_common_tests.rs index 7908f7f9b2..45d4334222 100644 --- a/codex-rs/core/src/client_common_tests.rs +++ b/codex-rs/core/src/client_common_tests.rs @@ -22,7 +22,6 @@ fn serializes_text_verbosity_when_set() { store: false, stream: true, include: vec![], - max_output_tokens: None, prompt_cache_key: None, service_tier: None, text: Some(TextControls { @@ -70,7 +69,6 @@ fn serializes_text_schema_with_strict_format() { store: false, stream: true, include: vec![], - max_output_tokens: None, prompt_cache_key: None, service_tier: None, text: Some(text_controls), @@ -132,7 +130,6 @@ fn omits_text_when_not_set() { store: false, stream: true, include: vec![], - max_output_tokens: None, prompt_cache_key: None, service_tier: None, text: None, @@ -156,7 +153,6 @@ fn serializes_flex_service_tier_when_set() { store: false, stream: true, include: vec![], - max_output_tokens: None, prompt_cache_key: None, service_tier: Some(ServiceTier::Flex.to_string()), text: None, diff --git a/codex-rs/core/src/next_prompt_suggestion.rs b/codex-rs/core/src/next_prompt_suggestion.rs index 13e2fa03a6..d3db54a1d4 100644 --- a/codex-rs/core/src/next_prompt_suggestion.rs +++ b/codex-rs/core/src/next_prompt_suggestion.rs @@ -73,9 +73,6 @@ pub(crate) async fn suggest_next_prompt( let started_at = Instant::now(); let mut turn_context = sess.new_lightweight_turn().await; prefer_fast_suggestion_profile(&mut turn_context); - if !suggestion_prompt_fits_context_window(sess, &turn_context).await { - return None; - } let history = sess.clone_history().await; let history_snapshot = HistorySnapshot::from_history(&history); @@ -111,6 +108,9 @@ pub(crate) async fn suggest_next_prompt( output_schema_strict: true, max_output_tokens: Some(NEXT_PROMPT_SUGGESTION_MAX_OUTPUT_TOKENS), }; + if !suggestion_prompt_fits_context_window(&prompt, &turn_context) { + return None; + } if !session_is_idle_for_suggestion(sess).await { return None; } @@ -276,14 +276,19 @@ fn history_matches_snapshot(history: &ContextManager, snapshot: HistorySnapshot) history.history_version() == snapshot.version && history.raw_items().len() == snapshot.len } -async fn suggestion_prompt_fits_context_window(sess: &Session, turn_context: &TurnContext) -> bool { +fn suggestion_prompt_fits_context_window(prompt: &Prompt, turn_context: &TurnContext) -> bool { let Some(model_context_window) = turn_context.model_context_window() else { tracing::debug!("next prompt suggestion skipped without model context window"); return false; }; - if let Some(estimated_token_count) = sess.get_estimated_token_count(turn_context).await - && !suggestion_prompt_has_headroom(estimated_token_count, model_context_window) - { + let Ok(input) = serde_json::to_string(&prompt.input) else { + tracing::debug!("next prompt suggestion skipped without serializable prompt input"); + return false; + }; + let estimated_token_count = approx_token_count(&prompt.base_instructions.text) + .saturating_add(approx_token_count(&input)); + let estimated_token_count = i64::try_from(estimated_token_count).unwrap_or(i64::MAX); + if !suggestion_prompt_has_headroom(estimated_token_count, model_context_window) { let suggestion_prompt_limit = model_context_window.saturating_sub(NEXT_PROMPT_SUGGESTION_TOKEN_HEADROOM); tracing::debug!( @@ -503,6 +508,7 @@ fn filter_next_prompt_suggestion(raw: &str) -> Option { ) || lower.starts_with("suggestion:") || lower.starts_with("next prompt:") || is_wrapped_meta(&suggestion) + || is_wrapped_quote(&suggestion) || starts_with_any(&lower, &["looks good", "thanks", "thank you"]) || starts_with_any(&lower, &["let me", "i'll", "i will", "here's"]) { @@ -524,6 +530,11 @@ fn is_wrapped_meta(suggestion: &str) -> bool { || (suggestion.starts_with('[') && suggestion.ends_with(']')) } +fn is_wrapped_quote(suggestion: &str) -> bool { + (suggestion.starts_with('"') && suggestion.ends_with('"')) + || (suggestion.starts_with('\'') && suggestion.ends_with('\'')) +} + fn starts_with_any(value: &str, prefixes: &[&str]) -> bool { prefixes.iter().any(|prefix| value.starts_with(prefix)) } diff --git a/codex-rs/core/src/next_prompt_suggestion_tests.rs b/codex-rs/core/src/next_prompt_suggestion_tests.rs index 8ddb265bd9..9adf250956 100644 --- a/codex-rs/core/src/next_prompt_suggestion_tests.rs +++ b/codex-rs/core/src/next_prompt_suggestion_tests.rs @@ -22,6 +22,7 @@ fn filter_keeps_valid_prompts() { ("set CODEX_HOME", "set CODEX_HOME"), ("update Cargo.toml", "update Cargo.toml"), ("open app-server/README.md", "open app-server/README.md"), + ("don't run tests yet", "don't run tests yet"), ] { assert_eq!( filter_next_prompt_suggestion(suggestion), @@ -235,6 +236,8 @@ fn filter_rejects_invalid_prompts() { "let me run tests", "what about tests?", "run tests.", + "\"run the tests\"", + "'run the tests'", "run\ntests", "continue with every possible next step in this project and explain every detail now", ] {