Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 25 additions & 11 deletions src/providers/codex/chat_completions/response.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
use http::StatusCode;
use serde_json::{Value, json};

use crate::anthropic::sse::parse_sse_events;
use crate::providers::codex::events::classify_event_failure;

use super::ChatError;

Expand Down Expand Up @@ -46,8 +48,16 @@ impl CompletionState {

pub fn observe(&mut self, event: &Value) -> Result<Option<String>, ChatError> {
let kind = event.get("type").and_then(Value::as_str);
if matches!(kind, Some("response.failed" | "response.error" | "error")) {
return Err(event_error(event));
if let Some(failure) = classify_event_failure(event) {
let status =
StatusCode::from_u16(failure.client_status()).unwrap_or(StatusCode::BAD_GATEWAY);
return Err(ChatError::new(
status,
failure.client_error_type(),
failure.message,
None,
None,
));
}
match kind {
Some("response.created" | "response.in_progress") => {
Expand All @@ -62,19 +72,14 @@ impl CompletionState {
self.text.push_str(delta);
return Ok(Some(delta.to_string()));
}
Some("response.completed" | "response.incomplete") => {
Some("response.completed" | "response.done") => {
let response = event.get("response").ok_or_else(|| {
ChatError::upstream("Codex completion event did not contain a response")
})?;
self.update_metadata(response);
if response.get("status").and_then(Value::as_str) == Some("failed") {
return Err(event_error(event));
}
if kind == Some("response.incomplete")
|| response.get("status").and_then(Value::as_str) == Some("incomplete")
{
self.finish_reason = "length";
}
self.completed = true;
}
_ => {
Expand Down Expand Up @@ -199,10 +204,19 @@ mod tests {
}

#[test]
fn incomplete_maps_to_length() {
fn incomplete_is_an_upstream_error() {
let body = b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"partial\"}\n\ndata: {\"type\":\"response.incomplete\",\"response\":{\"status\":\"incomplete\"}}\n\n";
let value = aggregate_sse(body, "model").unwrap();
assert_eq!(value["choices"][0]["finish_reason"], "length");
let error = aggregate_sse(body, "model").unwrap_err();
assert!(error.message.contains("Incomplete response"));
}

#[test]
fn buffered_chat_preserves_typed_terminal_failure() {
let body = b"data: {\"type\":\"response.failed\",\"response\":{\"status\":\"failed\",\"error\":{\"code\":\"context_length_exceeded\",\"message\":\"input exceeds context window\"}}}\n\n";
let error = aggregate_sse(body, "model").unwrap_err();
assert_eq!(error.status, StatusCode::PAYLOAD_TOO_LARGE);
assert_eq!(error.kind, "request_too_large");
assert_eq!(error.message, "input exceeds context window");
}

#[test]
Expand Down
222 changes: 130 additions & 92 deletions src/providers/codex/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1462,6 +1462,7 @@ impl CodexHttpClient {
);
}

let event_kind = super::events::classify_stream_event(&payload);
let failure = super::events::classify_event_failure(&payload);
if !semantic_output_forwarded
&& let Some(failure) = failure.as_ref()
Expand All @@ -1471,30 +1472,53 @@ impl CodexHttpClient {
break 'read_attempt codex_event_failure_error(failure.clone());
}

let terminal = event_closes_http_stream(&payload);
if !semantic_output_forwarded && failure.is_some() {
pending_events.clear();
if tx.send(Ok(payload)).await.is_err() {
return;
let terminal = super::events::event_is_terminal(&payload);
match event_kind {
super::events::CodexStreamEventKind::TerminalFailure => {
if !semantic_output_forwarded {
pending_events.clear();
}
if tx.send(Ok(payload)).await.is_err() {
return;
}
}
} else if !semantic_output_forwarded
&& http_event_starts_semantic_output(&payload)
{
semantic_output_forwarded = true;
for pending in pending_events.drain(..) {
if tx.send(Ok(pending)).await.is_err() {
super::events::CodexStreamEventKind::TerminalSuccess => {
for pending in pending_events.drain(..) {
if tx.send(Ok(pending)).await.is_err() {
return;
}
}
if tx.send(Ok(payload)).await.is_err() {
return;
}
}
if tx.send(Ok(payload)).await.is_err() {
return;
super::events::CodexStreamEventKind::Semantic => {
if !semantic_output_forwarded {
semantic_output_forwarded = true;
for pending in pending_events.drain(..) {
if tx.send(Ok(pending)).await.is_err() {
return;
}
}
}
if tx.send(Ok(payload)).await.is_err() {
return;
}
}
} else if semantic_output_forwarded || http_event_is_control(&payload) {
if tx.send(Ok(payload)).await.is_err() {
return;
super::events::CodexStreamEventKind::Control => {
if tx.send(Ok(payload)).await.is_err() {
return;
}
}
super::events::CodexStreamEventKind::Structural => {
if semantic_output_forwarded {
if tx.send(Ok(payload)).await.is_err() {
return;
}
} else {
pending_events.push(payload);
}
}
} else {
pending_events.push(payload);
}

if terminal {
Expand Down Expand Up @@ -2081,6 +2105,7 @@ impl CodexHttpClient {
}
};

let mut pending_events = Vec::new();
loop {
let item = tokio::select! {
biased;
Expand Down Expand Up @@ -2174,20 +2199,71 @@ impl CodexHttpClient {
continue 'attempt;
}

if item.as_ref().is_ok_and(event_closes_live_retry_window) {
forwarded_any = true;
}
let terminal = item.as_ref().is_err()
|| item.as_ref().is_ok_and(super::websocket::is_terminal_event);
if tx.send(item).await.is_err() {
if let Some(reservation) = continuation.as_ref() {
super::websocket::invalidate_codex_websocket_pool_socket(
reservation,
stream.socket_id(),
);
match item {
Ok(payload) => match super::events::classify_stream_event(&payload) {
super::events::CodexStreamEventKind::TerminalFailure => {
if !forwarded_any {
pending_events.clear();
}
if tx.send(Ok(payload)).await.is_err() {
return;
}
}
super::events::CodexStreamEventKind::TerminalSuccess => {
for pending in pending_events.drain(..) {
if tx.send(Ok(pending)).await.is_err() {
return;
}
}
if tx.send(Ok(payload)).await.is_err() {
return;
}
}
super::events::CodexStreamEventKind::Semantic => {
if !forwarded_any {
forwarded_any = true;
for pending in pending_events.drain(..) {
if tx.send(Ok(pending)).await.is_err() {
return;
}
}
}
if tx.send(Ok(payload)).await.is_err() {
return;
}
}
super::events::CodexStreamEventKind::Control => {
if tx.send(Ok(payload)).await.is_err() {
return;
}
}
super::events::CodexStreamEventKind::Structural => {
if forwarded_any {
if tx.send(Ok(payload)).await.is_err() {
return;
}
} else {
pending_events.push(payload);
}
}
},
Err(err) => {
if tx.send(Err(err)).await.is_err() {
if let Some(reservation) = continuation.as_ref() {
super::websocket::invalidate_codex_websocket_pool_socket(
reservation,
stream.socket_id(),
);
}
abort_abandoned_live_continuation(
continuation.as_ref(),
&socket_id_publisher,
);
return;
}
}
abort_abandoned_live_continuation(continuation.as_ref(), &socket_id_publisher);
return;
}
if terminal {
return;
Expand Down Expand Up @@ -2458,62 +2534,6 @@ fn response_headers(resp: &reqwest::Response) -> Vec<(String, String)> {
.collect()
}

fn event_closes_http_stream(payload: &serde_json::Value) -> bool {
matches!(
payload.get("type").and_then(|value| value.as_str()),
Some(
"response.completed"
| "response.incomplete"
| "response.done"
| "response.failed"
| "response.error"
| "error"
)
)
}

fn http_event_is_control(payload: &serde_json::Value) -> bool {
matches!(
payload.get("type").and_then(|value| value.as_str()),
Some(
"keepalive"
| "response.created"
| "response.in_progress"
| "codex.rate_limits"
| "response.web_search_call.in_progress"
| "response.web_search_call.searching"
| "response.web_search_call.completed"
)
)
}

fn http_event_starts_semantic_output(payload: &serde_json::Value) -> bool {
match payload.get("type").and_then(|value| value.as_str()) {
Some("response.output_item.added") => matches!(
payload
.pointer("/item/type")
.and_then(|value| value.as_str()),
Some("message" | "function_call")
),
Some(
"response.reasoning_summary_text.delta"
| "response.output_text.delta"
| "response.function_call_arguments.delta",
) => payload
.get("delta")
.and_then(|value| value.as_str())
.is_some_and(|delta| !delta.is_empty()),
Some("response.output_item.done") => matches!(
payload
.pointer("/item/type")
.and_then(|value| value.as_str()),
Some("reasoning" | "message" | "function_call")
),
Some("response.completed" | "response.incomplete" | "response.done") => true,
_ => false,
}
}

fn codex_event_failure_error(failure: super::events::CodexEventFailure) -> CodexError {
CodexError {
status: failure.status,
Expand Down Expand Up @@ -2953,11 +2973,9 @@ fn invalidate_live_continuation_pool(
super::websocket::invalidate_codex_websocket_pool_turn_for_owner(owner, continuation.turn_id());
}

#[cfg(test)]
fn event_closes_live_retry_window(payload: &serde_json::Value) -> bool {
!matches!(
payload.get("type").and_then(|value| value.as_str()),
Some("codex.rate_limits" | "keepalive")
)
super::events::classify_stream_event(payload) == super::events::CodexStreamEventKind::Semantic
}

pub(super) fn is_continuation_retry_error(err: &CodexError) -> bool {
Expand Down Expand Up @@ -3432,7 +3450,7 @@ mod tests {
.unwrap();
write_http_chunk(
&mut stream,
b"data: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"type\":\"message\",\"id\":\"msg_up\"}}\n\n",
b"data: {\"type\":\"response.output_text.delta\",\"output_index\":0,\"delta\":\"hello\"}\n\n",
)
.await;
release_rx.await.unwrap();
Expand Down Expand Up @@ -3462,7 +3480,7 @@ mod tests {
.unwrap();
assert_eq!(
first_upstream.get("type").and_then(|value| value.as_str()),
Some("response.output_item.added")
Some("response.output_text.delta")
);

release_tx.send(()).unwrap();
Expand Down Expand Up @@ -5244,9 +5262,29 @@ mod tests {
assert!(!event_closes_live_retry_window(&serde_json::json!({
"type": "keepalive"
})));
assert!(event_closes_live_retry_window(&serde_json::json!({
assert!(!event_closes_live_retry_window(&serde_json::json!({
"type": "response.created"
})));
assert!(!event_closes_live_retry_window(&serde_json::json!({
"type": "response.output_item.added",
"output_index": 0,
"item": {"type": "message", "id": "msg_1"}
})));
assert!(!event_closes_live_retry_window(&serde_json::json!({
"type": "response.output_item.added",
"output_index": 0,
"item": {"type": "function_call", "call_id": "call_1", "name": "Read"}
})));
assert!(event_closes_live_retry_window(&serde_json::json!({
"type": "response.output_text.delta",
"output_index": 0,
"delta": "hello"
})));
assert!(event_closes_live_retry_window(&serde_json::json!({
"type": "response.function_call_arguments.delta",
"output_index": 0,
"delta": "{}"
})));
}

#[test]
Expand Down
Loading