| 1 | //! Shared stream entry seam for Chat Completions / Anthropic Messages / Responses. |
| 2 | //! |
| 3 | //! Scoped consolidation for v0.9.1: wire-protocol adapters stay at the edge |
| 4 | //! (`chat.rs`, `anthropic.rs`, `responses.rs`); this module owns the common |
| 5 | //! open path, HTTP/1.1 fallback policy, and idle-timeout envelope so providers |
| 6 | //! do not re-implement transport differently. |
| 7 | //! |
| 8 | //! Full piagent-style provider collapse is deferred — see |
| 9 | //! `docs/notes/post-0.9.1-thin-tui-and-stream.md`. |
| 10 | |
| 11 | use std::future::Future; |
| 12 | use std::time::Duration; |
| 13 | |
| 14 | use anyhow::Result; |
| 15 | use reqwest::Client; |
| 16 | |
| 17 | use crate::llm_client::LlmError; |
| 18 | |
| 19 | /// Default bounded wait for SSE response headers. Intentionally shorter than |
| 20 | /// the per-chunk idle timeout: it covers connection setup and upstream header |
| 21 | /// return only, never model thinking time after streaming has started. |
| 22 | pub(crate) const DEFAULT_STREAM_OPEN_TIMEOUT: Duration = Duration::from_secs(45); |
| 23 | /// Accepted response-header wait range, in seconds. |
| 24 | pub(crate) const MIN_STREAM_OPEN_TIMEOUT_SECS: u64 = 5; |
| 25 | pub(crate) const MAX_STREAM_OPEN_TIMEOUT_SECS: u64 = 300; |
| 26 | |
| 27 | /// Resolve the response-header wait shared by every streaming adapter. |
| 28 | /// |
| 29 | /// A positive `[stream].open_timeout_secs` (legacy `[tui]` fallback) wins; omitted or `0` |
| 30 | /// falls back to the env override (`CODEWHALE_STREAM_OPEN_TIMEOUT_SECS`, |
| 31 | /// legacy `DEEPSEEK_STREAM_OPEN_TIMEOUT_SECS`), then the 45s default. Every |
| 32 | /// source clamps to `5..=300`. |
| 33 | #[must_use] |
| 34 | pub(crate) fn resolve_stream_open_timeout(configured_secs: Option<u64>) -> Duration { |
| 35 | match configured_secs { |
| 36 | Some(secs) if secs > 0 => Duration::from_secs( |
| 37 | secs.clamp(MIN_STREAM_OPEN_TIMEOUT_SECS, MAX_STREAM_OPEN_TIMEOUT_SECS), |
| 38 | ), |
| 39 | _ => stream_open_timeout_from_env( |
| 40 | std::env::var("CODEWHALE_STREAM_OPEN_TIMEOUT_SECS") |
| 41 | .or_else(|_| std::env::var("DEEPSEEK_STREAM_OPEN_TIMEOUT_SECS")) |
| 42 | .ok() |
| 43 | .as_deref(), |
| 44 | ), |
| 45 | } |
| 46 | } |
| 47 | |
| 48 | pub(crate) fn stream_open_timeout_from_env(value: Option<&str>) -> Duration { |
| 49 | let secs = value |
| 50 | .and_then(|v| v.parse::<u64>().ok()) |
| 51 | .unwrap_or(DEFAULT_STREAM_OPEN_TIMEOUT.as_secs()) |
| 52 | .clamp(MIN_STREAM_OPEN_TIMEOUT_SECS, MAX_STREAM_OPEN_TIMEOUT_SECS); |
| 53 | Duration::from_secs(secs) |
| 54 | } |
| 55 | |
| 56 | /// Default wait for the first body byte after the response headers (#6184). |
| 57 | /// Well under the 900s default inter-chunk idle budget: a provider that has |
| 58 | /// answered the headers and then sends nothing at all — not even an SSE |
| 59 | /// keep-alive — for five minutes has stopped, it is not thinking. Applies only |
| 60 | /// while the idle budget is the default; an explicitly configured |
| 61 | /// `stream_chunk_timeout_secs` is respected for the first byte too, so a user |
| 62 | /// who raised it for long silent reasoning keeps that allowance. |
| 63 | pub(crate) const DEFAULT_STREAM_FIRST_BYTE_TIMEOUT: Duration = Duration::from_secs(300); |
| 64 | |
| 65 | /// Resolve the first-body-byte bound for a stream whose inter-chunk idle |
| 66 | /// budget is `idle`. `CODEWHALE_STREAM_FIRST_BYTE_TIMEOUT_SECS` overrides it |
| 67 | /// (clamped to 5..=3600). |
| 68 | #[must_use] |
| 69 | pub(crate) fn first_byte_timeout(idle: Duration) -> Duration { |
| 70 | first_byte_timeout_from_env( |
| 71 | idle, |
| 72 | std::env::var("CODEWHALE_STREAM_FIRST_BYTE_TIMEOUT_SECS") |
| 73 | .ok() |
| 74 | .as_deref(), |
| 75 | ) |
| 76 | } |
| 77 | |
| 78 | pub(crate) fn first_byte_timeout_from_env(idle: Duration, value: Option<&str>) -> Duration { |
| 79 | if let Some(secs) = value.and_then(|v| v.trim().parse::<u64>().ok()) { |
| 80 | return Duration::from_secs(secs.clamp(5, 3600)); |
| 81 | } |
| 82 | let default_idle = Duration::from_secs(crate::config::DEFAULT_STREAM_CHUNK_TIMEOUT_SECS); |
| 83 | if idle == default_idle { |
| 84 | DEFAULT_STREAM_FIRST_BYTE_TIMEOUT.min(idle) |
| 85 | } else { |
| 86 | idle |
| 87 | } |
| 88 | } |
| 89 | |
| 90 | /// Bound for the next body read: the first-byte bound until any byte arrived, |
| 91 | /// the inter-chunk idle budget after. |
| 92 | #[must_use] |
| 93 | pub(crate) fn next_chunk_timeout( |
| 94 | idle: Duration, |
| 95 | first_byte: Duration, |
| 96 | bytes_received: usize, |
| 97 | ) -> Duration { |
| 98 | if bytes_received == 0 { |
| 99 | first_byte |
| 100 | } else { |
| 101 | idle |
| 102 | } |
| 103 | } |
| 104 | |
| 105 | /// Message and stall record for a body read that timed out. A first-byte |
| 106 | /// timeout is a stall worth a `crashes/` record (#6184); a later idle timeout |
| 107 | /// is reported the same way so every silent provider wait leaves a trace. |
| 108 | pub(crate) fn body_timeout_message( |
| 109 | timeout: Duration, |
| 110 | bytes_received: usize, |
| 111 | stream_age: Duration, |
| 112 | since_last_chunk: Duration, |
| 113 | provider: &str, |
| 114 | ) -> String { |
| 115 | let message = if bytes_received == 0 { |
| 116 | format!( |
| 117 | "SSE stream first-byte timeout after {}s — the provider sent headers but no data \ |
| 118 | (stream_age_ms={})", |
| 119 | timeout.as_secs(), |
| 120 | stream_age.as_millis(), |
| 121 | ) |
| 122 | } else { |
| 123 | idle_timeout_message(timeout, bytes_received, stream_age, since_last_chunk) |
| 124 | }; |
| 125 | let phase = if bytes_received == 0 { |
| 126 | "waiting for the provider's first byte" |
| 127 | } else { |
| 128 | "waiting for the next stream chunk" |
| 129 | }; |
| 130 | crate::core::engine::turn_heartbeat::report_stall( |
| 131 | &crate::core::engine::turn_heartbeat::StallReport { |
| 132 | source: "client", |
| 133 | phase: format!("while {phase}"), |
| 134 | detail: Some(provider.to_string()), |
| 135 | turn_id: None, |
| 136 | provider_request: None, |
| 137 | since_progress: since_last_chunk, |
| 138 | bound: Some(timeout), |
| 139 | }, |
| 140 | ); |
| 141 | message |
| 142 | } |
| 143 | |
| 144 | /// How the shared stream open path should pin HTTP version. |
| 145 | #[derive(Debug, Clone, Copy, PartialEq, Eq)] |
| 146 | pub enum StreamHttpPolicy { |
| 147 | /// Prefer the dual client (H2 primary, H1 twin for fallback). |
| 148 | DualWithH1Fallback, |
| 149 | /// Force HTTP/1.1 only (config/env pin or prior H2 stall). |
| 150 | Http1Only, |
| 151 | } |
| 152 | |
| 153 | /// Inputs shared by every streaming provider adapter at open time. |
| 154 | #[derive(Debug, Clone)] |
| 155 | pub struct StreamOpenRequest { |
| 156 | pub policy: StreamHttpPolicy, |
| 157 | pub open_timeout: Duration, |
| 158 | pub idle_timeout: Duration, |
| 159 | } |
| 160 | |
| 161 | impl StreamOpenRequest { |
| 162 | /// `force_http1` is the client's resolved pin (`Config::force_http1`: |
| 163 | /// `[stream].force_http1`, legacy `[tui]`, or `CODEWHALE_FORCE_HTTP1`), never re-read here. |
| 164 | #[must_use] |
| 165 | pub fn new(force_http1: bool, open_timeout: Duration, idle_timeout: Duration) -> Self { |
| 166 | Self { |
| 167 | policy: if force_http1 { |
| 168 | StreamHttpPolicy::Http1Only |
| 169 | } else { |
| 170 | StreamHttpPolicy::DualWithH1Fallback |
| 171 | }, |
| 172 | open_timeout, |
| 173 | idle_timeout, |
| 174 | } |
| 175 | } |
| 176 | |
| 177 | /// After an H2 stall, retry on the HTTP/1.1 twin. |
| 178 | #[must_use] |
| 179 | pub fn with_h1_only(mut self) -> Self { |
| 180 | self.policy = StreamHttpPolicy::Http1Only; |
| 181 | self |
| 182 | } |
| 183 | } |
| 184 | |
| 185 | impl super::CodewhaleClient { |
| 186 | /// The open request every streaming adapter starts from, carrying this |
| 187 | /// client's resolved HTTP/1.1 pin and timeouts (#6700). |
| 188 | #[must_use] |
| 189 | pub(super) fn stream_open_request(&self) -> StreamOpenRequest { |
| 190 | StreamOpenRequest::new( |
| 191 | self.force_http1, |
| 192 | self.stream_open_timeout, |
| 193 | self.stream_idle_timeout, |
| 194 | ) |
| 195 | } |
| 196 | } |
| 197 | |
| 198 | /// Select the HTTP client for a stream open attempt. |
| 199 | #[must_use] |
| 200 | pub fn client_for_policy<'a>( |
| 201 | primary: &'a Client, |
| 202 | http1_fallback: &'a Client, |
| 203 | policy: StreamHttpPolicy, |
| 204 | ) -> &'a Client { |
| 205 | match policy { |
| 206 | StreamHttpPolicy::DualWithH1Fallback => primary, |
| 207 | StreamHttpPolicy::Http1Only => http1_fallback, |
| 208 | } |
| 209 | } |
| 210 | |
| 211 | /// Whether a transport error should trigger H1 fallback retry. |
| 212 | #[must_use] |
| 213 | pub fn should_retry_with_h1(policy: StreamHttpPolicy, err_text: &str) -> bool { |
| 214 | if policy != StreamHttpPolicy::DualWithH1Fallback { |
| 215 | return false; |
| 216 | } |
| 217 | let lower = err_text.to_ascii_lowercase(); |
| 218 | lower.contains("http2") |
| 219 | || lower.contains("h2 ") |
| 220 | || lower.contains("stream closed") |
| 221 | || lower.contains("connection reset") |
| 222 | || lower.contains("protocol error") |
| 223 | || lower.contains("frame size") |
| 224 | } |
| 225 | |
| 226 | /// Whether an error raised before response headers should retry through the |
| 227 | /// HTTP/1.1 twin. Prefer typed transport errors, then retain the narrow text |
| 228 | /// classifier for lower-level H2 errors that reqwest exposes only as prose. |
| 229 | #[must_use] |
| 230 | fn should_retry_error_with_h1(policy: StreamHttpPolicy, err: &anyhow::Error) -> bool { |
| 231 | if policy != StreamHttpPolicy::DualWithH1Fallback { |
| 232 | return false; |
| 233 | } |
| 234 | |
| 235 | typed_open_transport_failure(err) |
| 236 | .unwrap_or_else(|| should_retry_with_h1(policy, &format!("{err:#}"))) |
| 237 | } |
| 238 | |
| 239 | /// Classify `err` by its types alone, searching the whole context chain: an |
| 240 | /// adapter's `.context("... request failed")` must not hide the transport |
| 241 | /// cause (#6711). `Some(true)` for a typed `LlmError::NetworkError`/`Timeout` |
| 242 | /// or a reqwest connect/timeout/request error that carries no HTTP status; |
| 243 | /// `Some(false)` for any other typed `LlmError` or reqwest error; `None` when |
| 244 | /// the chain holds neither type. |
| 245 | fn typed_open_transport_failure(err: &anyhow::Error) -> Option<bool> { |
| 246 | err.chain().find_map(|cause| { |
| 247 | if let Some(llm_error) = cause.downcast_ref::<LlmError>() { |
| 248 | return Some(matches!( |
| 249 | llm_error, |
| 250 | LlmError::NetworkError(_) | LlmError::Timeout(_) |
| 251 | )); |
| 252 | } |
| 253 | cause.downcast_ref::<reqwest::Error>().map(|reqwest_error| { |
| 254 | reqwest_error.status().is_none() |
| 255 | && (reqwest_error.is_connect() |
| 256 | || reqwest_error.is_timeout() |
| 257 | || reqwest_error.is_request()) |
| 258 | }) |
| 259 | }) |
| 260 | } |
| 261 | |
| 262 | /// Whether a failed stream open is a real transport failure: the request |
| 263 | /// never received response headers (#6699). Decided by type only: an untyped |
| 264 | /// error, such as an adapter's `HTTP 500 ...` rejection whose provider body |
| 265 | /// happens to mention a connection or timeout, is a provider answer and is |
| 266 | /// never classified as a transport failure here. |
| 267 | #[must_use] |
| 268 | pub(crate) fn is_stream_open_transport_failure(err: &anyhow::Error) -> bool { |
| 269 | typed_open_transport_failure(err).unwrap_or(false) |
| 270 | } |
| 271 | |
| 272 | /// Preserve provider-semantic failures returned by the H1 attempt. Only a |
| 273 | /// transport failure should be normalized into the shared retryable network |
| 274 | /// error; otherwise an auth or invalid-request failure could be retried as if |
| 275 | /// switching protocols had failed. |
| 276 | fn h1_fallback_error(err: anyhow::Error) -> anyhow::Error { |
| 277 | if let Some(llm_error) = err.downcast_ref::<LlmError>() |
| 278 | && !matches!(llm_error, LlmError::NetworkError(_) | LlmError::Timeout(_)) |
| 279 | { |
| 280 | return err; |
| 281 | } |
| 282 | |
| 283 | let detail = format!("{err:#}"); |
| 284 | if let Some(unreachable) = connect_failure_summary(&err) { |
| 285 | return anyhow::Error::new(LlmError::NetworkError(format!("{unreachable}: {detail}"))); |
| 286 | } |
| 287 | let mut message = format!("SSE stream request failed after HTTP/1.1 fallback: {detail}."); |
| 288 | // The HTTP/1.1 hint only helps when the protocol or TLS layer failed; it |
| 289 | // is noise for a refused connection or an unknown host. |
| 290 | if has_tls_cause(&err) || is_protocol_or_tls_failure(&detail) { |
| 291 | message.push_str( |
| 292 | " `codewhale doctor` can still pass when non-streaming requests work; \ |
| 293 | on Windows or proxy networks, try `CODEWHALE_FORCE_HTTP1=1` and rerun `codewhale`.", |
| 294 | ); |
| 295 | } |
| 296 | anyhow::Error::new(LlmError::NetworkError(message)) |
| 297 | } |
| 298 | |
| 299 | /// `Cannot reach host:port (<why>)` for a failure to open the connection. |
| 300 | /// reqwest also reports TLS handshake and certificate failures as connect |
| 301 | /// errors; those return `None` so they keep the HTTP/1.1 and proxy hint. |
| 302 | fn connect_failure_summary(err: &anyhow::Error) -> Option<String> { |
| 303 | let reqwest_error = err |
| 304 | .chain() |
| 305 | .find_map(|cause| cause.downcast_ref::<reqwest::Error>()) |
| 306 | .filter(|error| error.is_connect())?; |
| 307 | // Classify on the underlying causes only: reqwest's own message embeds |
| 308 | // the request URL, whose host or path could contain "dns" or "tls". |
| 309 | let mut causes = String::new(); |
| 310 | let mut source = std::error::Error::source(reqwest_error); |
| 311 | while let Some(cause) = source { |
| 312 | causes.push_str(&cause.to_string().to_ascii_lowercase()); |
| 313 | causes.push('\n'); |
| 314 | source = cause.source(); |
| 315 | } |
| 316 | if has_tls_cause(err) || is_protocol_or_tls_failure(&causes) { |
| 317 | return None; |
| 318 | } |
| 319 | let target = reqwest_error |
| 320 | .url() |
| 321 | .and_then(|url| { |
| 322 | Some(format!( |
| 323 | "{}:{}", |
| 324 | url.host_str()?, |
| 325 | url.port_or_known_default()? |
| 326 | )) |
| 327 | }) |
| 328 | .unwrap_or_else(|| "the provider host".to_string()); |
| 329 | let why = if causes.contains("refused") { |
| 330 | "connection refused" |
| 331 | } else if causes.contains("dns") || causes.contains("lookup") { |
| 332 | "DNS lookup failed" |
| 333 | } else { |
| 334 | "connection failed" |
| 335 | }; |
| 336 | // With a proxy configured, the host that refused may be the proxy. |
| 337 | let via_proxy = [ |
| 338 | "HTTPS_PROXY", |
| 339 | "https_proxy", |
| 340 | "ALL_PROXY", |
| 341 | "all_proxy", |
| 342 | "HTTP_PROXY", |
| 343 | "http_proxy", |
| 344 | ] |
| 345 | .iter() |
| 346 | .any(|name| std::env::var_os(name).is_some_and(|value| !value.is_empty())); |
| 347 | Some(if via_proxy { |
| 348 | format!("Cannot reach {target} or the configured proxy ({why})") |
| 349 | } else { |
| 350 | format!("Cannot reach {target} ({why})") |
| 351 | }) |
| 352 | } |
| 353 | |
| 354 | fn has_tls_cause(err: &anyhow::Error) -> bool { |
| 355 | err.chain().any(|mut cause| { |
| 356 | // hyper-rustls wraps tokio-rustls's IO error in another IO error; |
| 357 | // Error::source skips those inners, so inspect get_ref explicitly. |
| 358 | while let Some(inner) = cause |
| 359 | .downcast_ref::<std::io::Error>() |
| 360 | .and_then(std::io::Error::get_ref) |
| 361 | { |
| 362 | cause = inner; |
| 363 | } |
| 364 | cause.is::<rustls::Error>() |
| 365 | }) |
| 366 | } |
| 367 | |
| 368 | fn is_protocol_or_tls_failure(detail: &str) -> bool { |
| 369 | let lower = detail.to_ascii_lowercase(); |
| 370 | should_retry_with_h1(StreamHttpPolicy::DualWithH1Fallback, &lower) |
| 371 | || ["tls", "ssl", "certificate", "handshake", "alpn"] |
| 372 | .iter() |
| 373 | .any(|needle| lower.contains(needle)) |
| 374 | } |
| 375 | |
| 376 | /// Open an SSE response through the shared transport policy. |
| 377 | /// |
| 378 | /// `attempt` builds and sends one wire-specific request on the client |
| 379 | /// selected for the given policy (via [`client_for_policy`]); everything |
| 380 | /// transport-shared lives here: |
| 381 | /// |
| 382 | /// - the response-header wait is bounded by `open_req.open_timeout`; |
| 383 | /// - a classified transport failure or header stall on the dual client |
| 384 | /// retries exactly once on the HTTP/1.1 twin; |
| 385 | /// - a failure on an already H1-pinned request never retries; |
| 386 | /// - once response headers have been received the seam never retries — |
| 387 | /// body/stream errors belong to the adapter's decode loop. |
| 388 | pub(crate) async fn open_sse_response<F, Fut>( |
| 389 | open_req: &StreamOpenRequest, |
| 390 | attempt: F, |
| 391 | ) -> Result<reqwest::Response> |
| 392 | where |
| 393 | F: Fn(StreamHttpPolicy) -> Fut, |
| 394 | Fut: Future<Output = Result<reqwest::Response>>, |
| 395 | { |
| 396 | let fallback_reason = match tokio::time::timeout( |
| 397 | open_req.open_timeout, |
| 398 | attempt(open_req.policy), |
| 399 | ) |
| 400 | .await |
| 401 | { |
| 402 | Ok(Ok(response)) => return Ok(response), |
| 403 | Ok(Err(err)) => { |
| 404 | if !should_retry_error_with_h1(open_req.policy, &err) { |
| 405 | return Err(err); |
| 406 | } |
| 407 | "transport error before response headers" |
| 408 | } |
| 409 | Err(_elapsed) => { |
| 410 | if open_req.policy == StreamHttpPolicy::Http1Only { |
| 411 | return Err(anyhow::Error::new(LlmError::NetworkError(format!( |
| 412 | "SSE stream request did not receive response headers after {}s. \ |
| 413 | `codewhale doctor` can still pass when non-streaming requests work; \ |
| 414 | on Windows or proxy networks, try `CODEWHALE_FORCE_HTTP1=1` and rerun `codewhale`.", |
| 415 | open_req.open_timeout.as_secs() |
| 416 | )))); |
| 417 | } |
| 418 | "response-header timeout" |
| 419 | } |
| 420 | }; |
| 421 | |
| 422 | // No response body exists yet, so switching protocols and replaying the |
| 423 | // request cannot corrupt stream state. It can still bill twice if the |
| 424 | // provider accepted the first request (see the ambiguous-replay limit on |
| 425 | // `llm_client::with_retry`). The policy guard above keeps this to exactly |
| 426 | // one retry. |
| 427 | let h1_req = open_req.clone().with_h1_only(); |
| 428 | crate::logging::warn(format!( |
| 429 | "SSE stream {fallback_reason}; retrying once with HTTP/1.1" |
| 430 | )); |
| 431 | match tokio::time::timeout(h1_req.open_timeout, attempt(h1_req.policy)).await { |
| 432 | Ok(Ok(response)) => Ok(response), |
| 433 | Ok(Err(err)) => Err(h1_fallback_error(err)), |
| 434 | // Typed, not a bare string: a header stall is a transport |
| 435 | // failure, and `LlmError::NetworkError` is what the shared |
| 436 | // retry layer recognizes as retryable. As an untyped anyhow |
| 437 | // error this killed the whole turn outright. |
| 438 | Err(_elapsed) => Err(anyhow::Error::new(LlmError::NetworkError(format!( |
| 439 | "SSE stream request did not receive response headers after {}s \ |
| 440 | (HTTP/2 and HTTP/1.1). `codewhale doctor` can still pass when \ |
| 441 | non-streaming requests work; try `CODEWHALE_FORCE_HTTP1=1` and \ |
| 442 | rerun `codewhale`.", |
| 443 | open_req.open_timeout.as_secs() |
| 444 | )))), |
| 445 | } |
| 446 | } |
| 447 | |
| 448 | /// Format a stable idle-timeout message shared across adapters. |
| 449 | #[must_use] |
| 450 | pub fn idle_timeout_message( |
| 451 | idle: Duration, |
| 452 | bytes_received: usize, |
| 453 | stream_age: Duration, |
| 454 | since_last_chunk: Duration, |
| 455 | ) -> String { |
| 456 | format!( |
| 457 | "SSE stream idle timeout after {}s — no data received \ |
| 458 | (bytes_received={}, stream_age_ms={}, ms_since_last_chunk={})", |
| 459 | idle.as_secs(), |
| 460 | bytes_received, |
| 461 | stream_age.as_millis(), |
| 462 | since_last_chunk.as_millis(), |
| 463 | ) |
| 464 | } |
| 465 | |
| 466 | #[cfg(test)] |
| 467 | mod tests { |
| 468 | use std::sync::Arc; |
| 469 | use std::sync::atomic::{AtomicUsize, Ordering}; |
| 470 | |
| 471 | use wiremock::matchers::method; |
| 472 | use wiremock::{Mock, MockServer, ResponseTemplate}; |
| 473 | |
| 474 | use super::*; |
| 475 | |
| 476 | fn open_req(policy: StreamHttpPolicy, open_timeout: Duration) -> StreamOpenRequest { |
| 477 | StreamOpenRequest { |
| 478 | policy, |
| 479 | open_timeout, |
| 480 | idle_timeout: Duration::from_secs(30), |
| 481 | } |
| 482 | } |
| 483 | |
| 484 | async fn ok_server() -> MockServer { |
| 485 | let server = MockServer::start().await; |
| 486 | Mock::given(method("POST")) |
| 487 | .respond_with(ResponseTemplate::new(200)) |
| 488 | .mount(&server) |
| 489 | .await; |
| 490 | server |
| 491 | } |
| 492 | |
| 493 | #[test] |
| 494 | fn configured_http1_pin_controls_open_attempts_and_survives_clone() { |
| 495 | let _pin_env = super::super::tests::FORCE_HTTP1_ENV_LOCK.lock().unwrap(); |
| 496 | let _env = crate::test_support::lock_test_env(); |
| 497 | let _codewhale = crate::test_support::EnvVarGuard::remove("CODEWHALE_FORCE_HTTP1"); |
| 498 | let _deepseek = crate::test_support::EnvVarGuard::remove("DEEPSEEK_FORCE_HTTP1"); |
| 499 | let runtime = tokio::runtime::Builder::new_current_thread() |
| 500 | .enable_all() |
| 501 | .build() |
| 502 | .unwrap(); |
| 503 | for pinned in [true, false] { |
| 504 | let config: crate::config::Config = toml::from_str(&format!( |
| 505 | r#" |
| 506 | provider = "zai" |
| 507 | [providers.zai] |
| 508 | api_key = "stream-open-fixture-key" |
| 509 | [tui] |
| 510 | force_http1 = {pinned} |
| 511 | "# |
| 512 | )) |
| 513 | .unwrap(); |
| 514 | let client = super::super::CodewhaleClient::new(&config).unwrap(); |
| 515 | for client in [client.clone(), client] { |
| 516 | let attempts = AtomicUsize::new(0); |
| 517 | let result = |
| 518 | runtime.block_on(open_sse_response(&client.stream_open_request(), |policy| { |
| 519 | let attempt = attempts.fetch_add(1, Ordering::SeqCst); |
| 520 | async move { |
| 521 | if attempt == 0 { |
| 522 | // Only a genuinely dual-protocol request may |
| 523 | // retry this failure on its HTTP/1.1 twin. |
| 524 | return Err(anyhow::Error::new(LlmError::NetworkError( |
| 525 | "connection reset before response headers".into(), |
| 526 | ))); |
| 527 | } |
| 528 | assert_eq!(policy, StreamHttpPolicy::Http1Only); |
| 529 | Ok(reqwest::Response::from( |
| 530 | axum::http::Response::builder().status(200).body("")?, |
| 531 | )) |
| 532 | } |
| 533 | })); |
| 534 | if pinned { |
| 535 | assert!(result.is_err(), "config pin must forbid a second send"); |
| 536 | assert_eq!(attempts.load(Ordering::SeqCst), 1); |
| 537 | } else { |
| 538 | assert_eq!(result.unwrap().status(), 200); |
| 539 | assert_eq!(attempts.load(Ordering::SeqCst), 2); |
| 540 | } |
| 541 | } |
| 542 | } |
| 543 | } |
| 544 | |
| 545 | #[tokio::test] |
| 546 | async fn open_returns_first_attempt_response_on_dual_policy() { |
| 547 | let server = ok_server().await; |
| 548 | let client = crate::tls::reqwest_client(); |
| 549 | let attempts = Arc::new(AtomicUsize::new(0)); |
| 550 | let response = open_sse_response( |
| 551 | &open_req(StreamHttpPolicy::DualWithH1Fallback, Duration::from_secs(5)), |
| 552 | |policy| { |
| 553 | assert_eq!(policy, StreamHttpPolicy::DualWithH1Fallback); |
| 554 | let attempts = Arc::clone(&attempts); |
| 555 | let client = client.clone(); |
| 556 | let url = server.uri(); |
| 557 | async move { |
| 558 | attempts.fetch_add(1, Ordering::SeqCst); |
| 559 | Ok(client.post(url).send().await?) |
| 560 | } |
| 561 | }, |
| 562 | ) |
| 563 | .await |
| 564 | .expect("first attempt succeeds"); |
| 565 | assert_eq!(response.status(), 200); |
| 566 | assert_eq!(attempts.load(Ordering::SeqCst), 1); |
| 567 | } |
| 568 | |
| 569 | #[tokio::test] |
| 570 | async fn header_stall_on_dual_policy_retries_exactly_once_on_h1() { |
| 571 | let attempts = Arc::new(AtomicUsize::new(0)); |
| 572 | let response = open_sse_response( |
| 573 | &open_req( |
| 574 | StreamHttpPolicy::DualWithH1Fallback, |
| 575 | Duration::from_millis(150), |
| 576 | ), |
| 577 | |policy| { |
| 578 | let attempts = Arc::clone(&attempts); |
| 579 | async move { |
| 580 | let attempt = attempts.fetch_add(1, Ordering::SeqCst); |
| 581 | if attempt == 0 { |
| 582 | // First attempt stalls before response headers. |
| 583 | assert_eq!(policy, StreamHttpPolicy::DualWithH1Fallback); |
| 584 | std::future::pending::<()>().await; |
| 585 | } |
| 586 | assert_eq!(policy, StreamHttpPolicy::Http1Only); |
| 587 | // This test exercises retry policy, not loopback scheduling. |
| 588 | // A ready response keeps the 150 ms first-attempt timeout |
| 589 | // from also imposing a network deadline under suite load. |
| 590 | Ok(reqwest::Response::from( |
| 591 | axum::http::Response::builder().status(200).body("")?, |
| 592 | )) |
| 593 | } |
| 594 | }, |
| 595 | ) |
| 596 | .await |
| 597 | .expect("H1 fallback retry succeeds"); |
| 598 | assert_eq!(response.status(), 200); |
| 599 | assert_eq!( |
| 600 | attempts.load(Ordering::SeqCst), |
| 601 | 2, |
| 602 | "exactly one fallback retry" |
| 603 | ); |
| 604 | } |
| 605 | |
| 606 | #[tokio::test] |
| 607 | async fn transport_error_before_headers_retries_exactly_once_on_h1() { |
| 608 | let server = ok_server().await; |
| 609 | let client = crate::tls::reqwest_client(); |
| 610 | let attempts = Arc::new(AtomicUsize::new(0)); |
| 611 | let response = open_sse_response( |
| 612 | &open_req(StreamHttpPolicy::DualWithH1Fallback, Duration::from_secs(5)), |
| 613 | |policy| { |
| 614 | let attempts = Arc::clone(&attempts); |
| 615 | let client = client.clone(); |
| 616 | let url = server.uri(); |
| 617 | async move { |
| 618 | let attempt = attempts.fetch_add(1, Ordering::SeqCst); |
| 619 | if attempt == 0 { |
| 620 | assert_eq!(policy, StreamHttpPolicy::DualWithH1Fallback); |
| 621 | return Err(anyhow::Error::new(LlmError::NetworkError( |
| 622 | "connection reset before response headers".to_string(), |
| 623 | )) |
| 624 | .context("Chat API request failed")); |
| 625 | } |
| 626 | assert_eq!(policy, StreamHttpPolicy::Http1Only); |
| 627 | Ok(client.post(url).send().await?) |
| 628 | } |
| 629 | }, |
| 630 | ) |
| 631 | .await |
| 632 | .expect("H1 fallback retry succeeds after a transport error"); |
| 633 | assert_eq!(response.status(), 200); |
| 634 | assert_eq!( |
| 635 | attempts.load(Ordering::SeqCst), |
| 636 | 2, |
| 637 | "exactly one fallback retry" |
| 638 | ); |
| 639 | } |
| 640 | |
| 641 | #[tokio::test] |
| 642 | async fn header_stall_when_h1_pinned_never_retries_and_reports_timeout_text() { |
| 643 | let attempts = Arc::new(AtomicUsize::new(0)); |
| 644 | let err = open_sse_response( |
| 645 | &open_req(StreamHttpPolicy::Http1Only, Duration::from_millis(100)), |
| 646 | |_| { |
| 647 | let attempts = Arc::clone(&attempts); |
| 648 | async move { |
| 649 | attempts.fetch_add(1, Ordering::SeqCst); |
| 650 | std::future::pending::<()>().await; |
| 651 | unreachable!("stalled attempt never resolves") |
| 652 | } |
| 653 | }, |
| 654 | ) |
| 655 | .await |
| 656 | .expect_err("H1-pinned stall fails without retry"); |
| 657 | assert_eq!(attempts.load(Ordering::SeqCst), 1, "no retry when pinned"); |
| 658 | let text = err.to_string(); |
| 659 | assert!(text.contains("did not receive response headers"), "{text}"); |
| 660 | // The whole point of the typed error: a header stall must reach the |
| 661 | // shared retry layer as retryable. As an untyped anyhow error it |
| 662 | // killed the turn outright, so a long root run lost all its work while |
| 663 | // a sub-agent — which text-matches the same message in its own |
| 664 | // classifier — would have retried and continued. |
| 665 | let classified = err |
| 666 | .downcast_ref::<crate::llm_client::LlmError>() |
| 667 | .expect("header stall must be a typed LlmError"); |
| 668 | assert!( |
| 669 | classified.is_retryable(), |
| 670 | "header stall must be retryable: {classified:?}" |
| 671 | ); |
| 672 | assert!(text.contains("CODEWHALE_FORCE_HTTP1=1"), "{text}"); |
| 673 | assert!( |
| 674 | !text.contains("HTTP/2 and HTTP/1.1"), |
| 675 | "single-protocol stall must not claim a dual-protocol attempt: {text}" |
| 676 | ); |
| 677 | } |
| 678 | |
| 679 | #[tokio::test] |
| 680 | async fn transport_error_when_h1_pinned_never_retries() { |
| 681 | let attempts = Arc::new(AtomicUsize::new(0)); |
| 682 | let err = open_sse_response( |
| 683 | &open_req(StreamHttpPolicy::Http1Only, Duration::from_secs(5)), |
| 684 | |_| { |
| 685 | let attempts = Arc::clone(&attempts); |
| 686 | async move { |
| 687 | attempts.fetch_add(1, Ordering::SeqCst); |
| 688 | Err(anyhow::Error::new(LlmError::NetworkError( |
| 689 | "connection reset before response headers".to_string(), |
| 690 | ))) |
| 691 | } |
| 692 | }, |
| 693 | ) |
| 694 | .await |
| 695 | .expect_err("an H1-pinned transport error must not retry"); |
| 696 | assert_eq!(attempts.load(Ordering::SeqCst), 1, "no retry when pinned"); |
| 697 | assert!(err.to_string().contains("connection reset"), "{err}"); |
| 698 | } |
| 699 | |
| 700 | #[tokio::test] |
| 701 | async fn refused_connection_names_the_host_and_skips_the_http1_hint() { |
| 702 | let port = { |
| 703 | let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind"); |
| 704 | listener.local_addr().expect("addr").port() |
| 705 | }; |
| 706 | let client = crate::tls::reqwest_client(); |
| 707 | let url = format!("http://127.0.0.1:{port}/v1/chat"); |
| 708 | let err = open_sse_response( |
| 709 | &open_req(StreamHttpPolicy::DualWithH1Fallback, Duration::from_secs(5)), |
| 710 | |_| { |
| 711 | let request = client.post(&url); |
| 712 | async move { Ok::<_, anyhow::Error>(request.send().await?) } |
| 713 | }, |
| 714 | ) |
| 715 | .await |
| 716 | .expect_err("nothing listens on the port"); |
| 717 | let text = err.to_string(); |
| 718 | assert!( |
| 719 | text.contains(&format!("Cannot reach 127.0.0.1:{port}")) |
| 720 | && text.contains("(connection refused)"), |
| 721 | "{text}" |
| 722 | ); |
| 723 | assert!(!text.contains("CODEWHALE_FORCE_HTTP1"), "{text}"); |
| 724 | assert!( |
| 725 | err.downcast_ref::<LlmError>() |
| 726 | .is_some_and(LlmError::is_retryable), |
| 727 | "{err:#}" |
| 728 | ); |
| 729 | } |
| 730 | |
| 731 | #[tokio::test] |
| 732 | async fn tls_failure_during_connect_keeps_the_http1_hint() { |
| 733 | // A plain-TCP peer on an https URL fails the TLS handshake, which |
| 734 | // reqwest reports as a connect error; it must not read as |
| 735 | // "Cannot reach" and must keep the FORCE_HTTP1/proxy hint. |
| 736 | let listener = tokio::net::TcpListener::bind("127.0.0.1:0") |
| 737 | .await |
| 738 | .expect("bind"); |
| 739 | let port = listener.local_addr().expect("addr").port(); |
| 740 | let server = tokio::spawn(async move { |
| 741 | use tokio::io::AsyncWriteExt; |
| 742 | while let Ok((mut socket, _)) = listener.accept().await { |
| 743 | let _ = socket.write_all(b"HTTP/1.1 400 Not TLS\r\n\r\n").await; |
| 744 | // Keep reading until the client closes. Dropping a socket with |
| 745 | // unread ClientHello bytes can send a TCP reset on Windows, |
| 746 | // hiding the TLS protocol error this fixture is meant to test. |
| 747 | let _ = tokio::io::copy(&mut socket, &mut tokio::io::sink()).await; |
| 748 | } |
| 749 | }); |
| 750 | crate::tls::ensure_rustls_crypto_provider(); |
| 751 | let client = crate::tls::reqwest_client(); |
| 752 | let url = format!("https://127.0.0.1:{port}/v1/chat"); |
| 753 | let err = open_sse_response( |
| 754 | &open_req(StreamHttpPolicy::DualWithH1Fallback, Duration::from_secs(5)), |
| 755 | |_| { |
| 756 | let request = client.post(&url); |
| 757 | async move { Ok::<_, anyhow::Error>(request.send().await?) } |
| 758 | }, |
| 759 | ) |
| 760 | .await |
| 761 | .expect_err("the peer does not speak TLS"); |
| 762 | server.abort(); |
| 763 | let text = err.to_string(); |
| 764 | assert!(!text.contains("Cannot reach"), "{text}"); |
| 765 | assert!(text.contains("CODEWHALE_FORCE_HTTP1=1"), "{text}"); |
| 766 | } |
| 767 | |
| 768 | #[test] |
| 769 | fn h1_fallback_hint_follows_protocol_errors_with_the_full_chain() { |
| 770 | let protocol = h1_fallback_error( |
| 771 | anyhow::Error::new(LlmError::NetworkError("http2 protocol error".to_string())) |
| 772 | .context("sending stream request"), |
| 773 | ); |
| 774 | let text = protocol.to_string(); |
| 775 | assert!(text.contains("sending stream request: "), "{text}"); |
| 776 | assert!(text.contains("http2 protocol error"), "{text}"); |
| 777 | assert!(text.contains("CODEWHALE_FORCE_HTTP1=1"), "{text}"); |
| 778 | |
| 779 | let other = h1_fallback_error(anyhow::Error::new(LlmError::NetworkError( |
| 780 | "unexpected end of body".to_string(), |
| 781 | ))); |
| 782 | assert!( |
| 783 | !other.to_string().contains("CODEWHALE_FORCE_HTTP1"), |
| 784 | "{other}" |
| 785 | ); |
| 786 | } |
| 787 | |
| 788 | #[tokio::test] |
| 789 | async fn provider_error_from_h1_fallback_keeps_its_semantic_type() { |
| 790 | let attempts = Arc::new(AtomicUsize::new(0)); |
| 791 | let err = open_sse_response( |
| 792 | &open_req(StreamHttpPolicy::DualWithH1Fallback, Duration::from_secs(5)), |
| 793 | |policy| { |
| 794 | let attempts = Arc::clone(&attempts); |
| 795 | async move { |
| 796 | let attempt = attempts.fetch_add(1, Ordering::SeqCst); |
| 797 | if attempt == 0 { |
| 798 | assert_eq!(policy, StreamHttpPolicy::DualWithH1Fallback); |
| 799 | return Err(anyhow::Error::new(LlmError::NetworkError( |
| 800 | "connection reset before response headers".to_string(), |
| 801 | ))); |
| 802 | } |
| 803 | assert_eq!(policy, StreamHttpPolicy::Http1Only); |
| 804 | Err(anyhow::Error::new(LlmError::InvalidRequest { |
| 805 | status: 400, |
| 806 | message: "invalid request".to_string(), |
| 807 | })) |
| 808 | } |
| 809 | }, |
| 810 | ) |
| 811 | .await |
| 812 | .expect_err("the H1 provider error must be returned"); |
| 813 | assert_eq!(attempts.load(Ordering::SeqCst), 2, "one fallback attempt"); |
| 814 | assert!( |
| 815 | matches!( |
| 816 | err.downcast_ref::<LlmError>(), |
| 817 | Some(LlmError::InvalidRequest { status: 400, .. }) |
| 818 | ), |
| 819 | "provider error was reclassified: {err:#}" |
| 820 | ); |
| 821 | } |
| 822 | |
| 823 | #[tokio::test] |
| 824 | async fn attempt_error_before_headers_is_not_h1_retried() { |
| 825 | let attempts = Arc::new(AtomicUsize::new(0)); |
| 826 | let err = open_sse_response( |
| 827 | &open_req(StreamHttpPolicy::DualWithH1Fallback, Duration::from_secs(5)), |
| 828 | |_| { |
| 829 | let attempts = Arc::clone(&attempts); |
| 830 | async move { |
| 831 | attempts.fetch_add(1, Ordering::SeqCst); |
| 832 | Err(anyhow::anyhow!("HTTP 401: invalid api key")) |
| 833 | } |
| 834 | }, |
| 835 | ) |
| 836 | .await |
| 837 | .expect_err("provider error propagates"); |
| 838 | assert_eq!( |
| 839 | attempts.load(Ordering::SeqCst), |
| 840 | 1, |
| 841 | "non-stall errors are never H1-retried" |
| 842 | ); |
| 843 | assert!(err.to_string().contains("HTTP 401"), "{err}"); |
| 844 | } |
| 845 | |
| 846 | #[tokio::test] |
| 847 | async fn double_stall_reports_both_protocols_in_timeout_text() { |
| 848 | let attempts = Arc::new(AtomicUsize::new(0)); |
| 849 | let err = open_sse_response( |
| 850 | &open_req( |
| 851 | StreamHttpPolicy::DualWithH1Fallback, |
| 852 | Duration::from_millis(100), |
| 853 | ), |
| 854 | |_| { |
| 855 | let attempts = Arc::clone(&attempts); |
| 856 | async move { |
| 857 | attempts.fetch_add(1, Ordering::SeqCst); |
| 858 | std::future::pending::<()>().await; |
| 859 | unreachable!("stalled attempt never resolves") |
| 860 | } |
| 861 | }, |
| 862 | ) |
| 863 | .await |
| 864 | .expect_err("double stall fails"); |
| 865 | assert_eq!(attempts.load(Ordering::SeqCst), 2, "one fallback, no more"); |
| 866 | let text = err.to_string(); |
| 867 | assert!(text.contains("HTTP/2 and HTTP/1.1"), "{text}"); |
| 868 | } |
| 869 | |
| 870 | #[test] |
| 871 | fn h1_retry_only_on_dual_policy() { |
| 872 | assert!(should_retry_with_h1( |
| 873 | StreamHttpPolicy::DualWithH1Fallback, |
| 874 | "http2 protocol error" |
| 875 | )); |
| 876 | assert!(!should_retry_with_h1( |
| 877 | StreamHttpPolicy::Http1Only, |
| 878 | "http2 protocol error" |
| 879 | )); |
| 880 | } |
| 881 | |
| 882 | #[test] |
| 883 | fn stall_first_byte_timeout_is_well_under_default_idle_budget() { |
| 884 | let default_idle = Duration::from_secs(crate::config::DEFAULT_STREAM_CHUNK_TIMEOUT_SECS); |
| 885 | let first_byte = first_byte_timeout_from_env(default_idle, None); |
| 886 | assert_eq!(first_byte, DEFAULT_STREAM_FIRST_BYTE_TIMEOUT); |
| 887 | assert!( |
| 888 | first_byte * 3 <= default_idle, |
| 889 | "{first_byte:?} vs {default_idle:?}" |
| 890 | ); |
| 891 | // An explicitly configured idle budget is respected for the first byte. |
| 892 | let custom = Duration::from_secs(1800); |
| 893 | assert_eq!(first_byte_timeout_from_env(custom, None), custom); |
| 894 | assert_eq!( |
| 895 | first_byte_timeout_from_env(Duration::from_secs(60), None), |
| 896 | Duration::from_secs(60) |
| 897 | ); |
| 898 | assert_eq!( |
| 899 | first_byte_timeout_from_env(default_idle, Some("90")), |
| 900 | Duration::from_secs(90) |
| 901 | ); |
| 902 | assert_eq!(next_chunk_timeout(default_idle, first_byte, 0), first_byte); |
| 903 | assert_eq!( |
| 904 | next_chunk_timeout(default_idle, first_byte, 1), |
| 905 | default_idle |
| 906 | ); |
| 907 | } |
| 908 | |
| 909 | #[test] |
| 910 | fn stream_open_timeout_defaults_and_clamps_env_values() { |
| 911 | assert_eq!(stream_open_timeout_from_env(None), Duration::from_secs(45)); |
| 912 | assert_eq!( |
| 913 | stream_open_timeout_from_env(Some("not-a-number")), |
| 914 | Duration::from_secs(45) |
| 915 | ); |
| 916 | assert_eq!( |
| 917 | stream_open_timeout_from_env(Some("1")), |
| 918 | Duration::from_secs(5) |
| 919 | ); |
| 920 | assert_eq!( |
| 921 | stream_open_timeout_from_env(Some("120")), |
| 922 | Duration::from_secs(120) |
| 923 | ); |
| 924 | assert_eq!( |
| 925 | stream_open_timeout_from_env(Some("999")), |
| 926 | Duration::from_secs(300) |
| 927 | ); |
| 928 | } |
| 929 | |
| 930 | #[test] |
| 931 | fn configured_stream_open_timeout_wins_and_clamps() { |
| 932 | // Positive config values never consult the env, so these are |
| 933 | // deterministic regardless of the test process environment. |
| 934 | assert_eq!( |
| 935 | resolve_stream_open_timeout(Some(90)), |
| 936 | Duration::from_secs(90) |
| 937 | ); |
| 938 | assert_eq!( |
| 939 | resolve_stream_open_timeout(Some(1)), |
| 940 | Duration::from_secs(MIN_STREAM_OPEN_TIMEOUT_SECS) |
| 941 | ); |
| 942 | assert_eq!( |
| 943 | resolve_stream_open_timeout(Some(u64::MAX)), |
| 944 | Duration::from_secs(MAX_STREAM_OPEN_TIMEOUT_SECS) |
| 945 | ); |
| 946 | } |
| 947 | |
| 948 | #[test] |
| 949 | fn idle_message_is_stable() { |
| 950 | let msg = idle_timeout_message( |
| 951 | Duration::from_secs(30), |
| 952 | 0, |
| 953 | Duration::from_secs(30), |
| 954 | Duration::from_secs(30), |
| 955 | ); |
| 956 | assert!(msg.contains("idle timeout")); |
| 957 | assert!(msg.contains("bytes_received=0")); |
| 958 | } |
| 959 | } |
| 960 |