use std::{error::Error, fmt};
use crate::providers::CODEX_RESPONSES_URL;
pub(crate) const CODEX_SESSION_EXPIRED_MESSAGE: &str =
"Your ChatGPT session expired before this request finished.";
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
pub(crate) struct ProviderStreamTrace {
pub(crate) schema_version: u64,
pub(crate) provider: String,
pub(crate) failure_context: String,
pub(crate) message_delta_stop_reason: Option<String>,
pub(crate) recent_events: Vec<ProviderStreamTraceEvent>,
pub(crate) pending_tool_count: usize,
pub(crate) pending_tools_truncated: bool,
pub(crate) pending_tools: Vec<ProviderStreamTracePendingTool>,
}
const RESPONSE_IDENTITY_STRING_MAX_BYTES: usize = 256;
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub(crate) struct ResponseAttemptIdentity {
pub(crate) schema_version: u64,
pub(crate) provider: String,
pub(crate) attempt: usize,
pub(crate) requested_model: String,
pub(crate) provider_response_model: Option<String>,
pub(crate) request_id: Option<String>,
pub(crate) outcome: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub(crate) cache: Option<Box<super::codex_session::CacheDiagnostics>>,
}
pub(crate) fn bounded_response_identity_string(value: &str) -> String {
let value = crate::output::redact_sensitive_text(value);
let end = value
.as_bytes()
.get(..RESPONSE_IDENTITY_STRING_MAX_BYTES)
.map_or(value.len(), |prefix| {
let mut end = prefix.len();
while end > 0 && !value.is_char_boundary(end) {
end -= 1;
}
end
});
value[..end].to_string()
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
pub(crate) struct ProviderStreamTraceEvent {
pub(crate) seq: u64,
pub(crate) event_type: String,
pub(crate) index: Option<u64>,
pub(crate) content_block_type: Option<String>,
pub(crate) delta_type: Option<String>,
pub(crate) message_delta_stop_reason: Option<String>,
pub(crate) usage: Option<ProviderStreamTraceUsage>,
pub(crate) partial_json_bytes: Option<usize>,
pub(crate) partial_json_sha256: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
pub(crate) struct ProviderStreamTraceUsage {
pub(crate) input_tokens: Option<u64>,
pub(crate) output_tokens: Option<u64>,
pub(crate) cache_read_input_tokens: Option<u64>,
pub(crate) cache_creation_input_tokens: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
pub(crate) struct ProviderStreamTracePendingTool {
pub(crate) index: u64,
pub(crate) id: String,
pub(crate) name: String,
pub(crate) argument_bytes: usize,
pub(crate) argument_sha256: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum ProviderErrorKind {
HttpStatus { status: u16 },
PromptTooLong,
Transport,
StreamTerminal,
StreamFailedIncomplete,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct ProviderError {
kind: ProviderErrorKind,
message: String,
stream_trace: Option<ProviderStreamTrace>,
}
impl ProviderError {
pub(crate) fn prompt_too_long(message: impl Into<String>) -> Self {
Self {
kind: ProviderErrorKind::PromptTooLong,
message: message.into(),
stream_trace: None,
}
}
pub(crate) fn is_prompt_too_long(&self) -> bool {
self.kind == ProviderErrorKind::PromptTooLong
}
pub(crate) fn http_status(status: u16, message: impl Into<String>) -> Self {
Self {
kind: ProviderErrorKind::HttpStatus { status },
message: message.into(),
stream_trace: None,
}
}
pub(crate) fn transport(message: impl Into<String>) -> Self {
Self {
kind: ProviderErrorKind::Transport,
message: message.into(),
stream_trace: None,
}
}
pub(crate) fn stream_terminal(message: impl Into<String>) -> Self {
Self {
kind: ProviderErrorKind::StreamTerminal,
message: message.into(),
stream_trace: None,
}
}
pub(crate) fn stream_failed_incomplete(message: impl Into<String>) -> Self {
Self {
kind: ProviderErrorKind::StreamFailedIncomplete,
message: message.into(),
stream_trace: None,
}
}
pub(crate) fn stream_trace(&self) -> Option<&ProviderStreamTrace> {
self.stream_trace.as_ref()
}
pub(crate) fn with_stream_trace(mut self, trace: ProviderStreamTrace) -> Self {
self.stream_trace = Some(trace);
self
}
pub(crate) fn is_stream_failed_incomplete(&self) -> bool {
self.kind == ProviderErrorKind::StreamFailedIncomplete
}
pub(crate) fn http_status_code(&self) -> Option<u16> {
match self.kind {
ProviderErrorKind::HttpStatus { status } => Some(status),
ProviderErrorKind::Transport
| ProviderErrorKind::PromptTooLong
| ProviderErrorKind::StreamTerminal
| ProviderErrorKind::StreamFailedIncomplete => None,
}
}
pub(crate) fn is_codex_session_expired_401(&self) -> bool {
self.http_status_code() == Some(401)
&& self.message.contains(CODEX_RESPONSES_URL)
&& self.message.contains(CODEX_SESSION_EXPIRED_MESSAGE)
}
pub(crate) fn is_retryable(&self) -> bool {
match self.kind {
ProviderErrorKind::PromptTooLong => false,
ProviderErrorKind::HttpStatus { status } => {
matches!(status, 408 | 425 | 429 | 500 | 502 | 503 | 504)
}
ProviderErrorKind::Transport
| ProviderErrorKind::StreamTerminal
| ProviderErrorKind::StreamFailedIncomplete => true,
}
}
}
impl fmt::Display for ProviderError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(&self.message)
}
}
impl Error for ProviderError {}
pub(crate) fn retryable_provider_error(error: &anyhow::Error) -> bool {
error.chain().any(|cause| {
cause
.downcast_ref::<ProviderError>()
.is_some_and(ProviderError::is_retryable)
})
}
pub(crate) fn codex_session_expired_provider_error(error: &anyhow::Error) -> bool {
error
.downcast_ref::<ProviderError>()
.is_some_and(ProviderError::is_codex_session_expired_401)
}
pub(crate) fn incomplete_stream_error(error: &anyhow::Error) -> bool {
error.chain().any(|cause| {
cause
.downcast_ref::<ProviderError>()
.is_some_and(ProviderError::is_stream_failed_incomplete)
})
}
pub(crate) fn provider_stream_trace_from_error(
error: &anyhow::Error,
) -> Option<ProviderStreamTrace> {
error.chain().find_map(|cause| {
cause
.downcast_ref::<ProviderError>()
.and_then(ProviderError::stream_trace)
.cloned()
})
}