diff --git a/codex-rs/codex-api/src/auth.rs b/codex-rs/codex-api/src/auth.rs index 41394a2258..a6fd6622f6 100644 --- a/codex-rs/codex-api/src/auth.rs +++ b/codex-rs/codex-api/src/auth.rs @@ -41,6 +41,27 @@ pub trait AuthProvider: Send + Sync { headers } + /// Adds auth headers for a concrete HTTP request URL. + /// + /// URL-sensitive providers can override this to scope request-specific + /// headers without exposing them through generic/non-HTTP auth helpers. + fn add_auth_headers_for_url(&self, _request_url: &str, headers: &mut HeaderMap) { + self.add_auth_headers(headers); + } + + /// Observes response headers for auth state that may need to rotate after a request. + /// + /// Most providers do not need this. Providers with server-minted, + /// request-scoped state may validate the URL and selectively persist + /// response headers here. + fn observe_response_headers( + &self, + _request_url: &str, + _request_headers: &HeaderMap, + _response_headers: &HeaderMap, + ) { + } + /// Applies auth to a complete outbound request and returns the request to send. /// /// The input `request` is moved into this method. Implementations may mutate @@ -54,7 +75,7 @@ pub trait AuthProvider: Send + Sync { /// If this returns [`AuthError`], the request should not be sent. async fn apply_auth(&self, request: Request) -> Result { let mut request = request; - self.add_auth_headers(&mut request.headers); + self.add_auth_headers_for_url(&request.url, &mut request.headers); Ok(request) } } diff --git a/codex-rs/codex-api/src/endpoint/responses_websocket.rs b/codex-rs/codex-api/src/endpoint/responses_websocket.rs index f0a0019817..022f2fd0b8 100644 --- a/codex-rs/codex-api/src/endpoint/responses_websocket.rs +++ b/codex-rs/codex-api/src/endpoint/responses_websocket.rs @@ -377,13 +377,33 @@ impl ResponsesWebsocketClient { .provider .websocket_url_for_path("responses") .map_err(|err| ApiError::Stream(format!("failed to build websocket URL: {err}")))?; + let request_url = ws_url.to_string(); let mut headers = merge_request_headers(&self.provider.headers, extra_headers, default_headers); - self.auth.add_auth_headers(&mut headers); + self.auth + .add_auth_headers_for_url(&request_url, &mut headers); + let request_headers = headers.clone(); - let (stream, _status, server_reasoning_included, models_etag, server_model) = - connect_websocket(ws_url, headers, turn_state.clone()).await?; + let ( + stream, + _status, + server_reasoning_included, + models_etag, + server_model, + response_headers, + ) = connect_websocket(ws_url, headers, turn_state.clone()) + .await + .inspect_err(|error| { + observe_websocket_auth_error_headers( + self.auth.as_ref(), + &request_url, + &request_headers, + error, + ); + })?; + self.auth + .observe_response_headers(&request_url, &request_headers, &response_headers); Ok(ResponsesWebsocketConnection::new( stream, self.provider.stream_idle_timeout, @@ -411,13 +431,27 @@ impl ResponsesWebsocketClient { .provider .websocket_url_for_path("responses") .map_err(|err| ApiError::Stream(format!("failed to build websocket URL: {err}")))?; + let request_url = ws_url.to_string(); let mut headers = merge_request_headers(&self.provider.headers, extra_headers, default_headers); - self.auth.add_auth_headers(&mut headers); + self.auth + .add_auth_headers_for_url(&request_url, &mut headers); + let request_headers = headers.clone(); - let (mut stream, status, reasoning_included, models_etag, server_model) = - connect_websocket(ws_url.clone(), headers, /*turn_state*/ None).await?; + let (mut stream, status, reasoning_included, models_etag, server_model, response_headers) = + connect_websocket(ws_url.clone(), headers, /*turn_state*/ None) + .await + .inspect_err(|error| { + observe_websocket_auth_error_headers( + self.auth.as_ref(), + &request_url, + &request_headers, + error, + ); + })?; + self.auth + .observe_response_headers(&request_url, &request_headers, &response_headers); let immediate_close = tokio::time::timeout(immediate_close_timeout, stream.next()) .await .ok() @@ -439,6 +473,21 @@ impl ResponsesWebsocketClient { } } +fn observe_websocket_auth_error_headers( + auth: &dyn crate::auth::AuthProvider, + request_url: &str, + request_headers: &HeaderMap, + error: &ApiError, +) { + if let ApiError::Transport(TransportError::Http { + headers: Some(response_headers), + .. + }) = error + { + auth.observe_response_headers(request_url, request_headers, response_headers); + } +} + fn immediate_close_from_message(message: Message) -> Option { let Message::Close(frame) = message else { return None; @@ -472,7 +521,17 @@ async fn connect_websocket( url: Url, headers: HeaderMap, turn_state: Option>>, -) -> Result<(WsStream, StatusCode, bool, Option, Option), ApiError> { +) -> Result< + ( + WsStream, + StatusCode, + bool, + Option, + Option, + HeaderMap, + ), + ApiError, +> { ensure_rustls_crypto_provider(); info!("connecting to websocket: {url}"); @@ -499,10 +558,7 @@ async fn connect_websocket( let (stream, response) = match response { Ok((stream, response)) => { - info!( - "successfully connected to websocket: {url}, headers: {:?}", - response.headers() - ); + info!("successfully connected to websocket: {url}"); (stream, response) } Err(err) => { @@ -536,6 +592,7 @@ async fn connect_websocket( reasoning_included, models_etag, server_model, + response.headers().clone(), )) } diff --git a/codex-rs/codex-api/src/endpoint/session.rs b/codex-rs/codex-api/src/endpoint/session.rs index 132c3abd90..528f559dc7 100644 --- a/codex-rs/codex-api/src/endpoint/session.rs +++ b/codex-rs/codex-api/src/endpoint/session.rs @@ -102,7 +102,16 @@ impl EndpointSession { let transport = &self.transport; async move { let req = auth.apply_auth(req).await.map_err(TransportError::from)?; - transport.execute(req).await + let request_url = req.url.clone(); + let request_headers = req.headers.clone(); + let response = transport.execute(req).await; + observe_auth_response_headers( + auth.as_ref(), + &request_url, + &request_headers, + &response, + ); + response } }, ) @@ -143,7 +152,16 @@ impl EndpointSession { let transport = &self.transport; async move { let req = auth.apply_auth(req).await.map_err(TransportError::from)?; - transport.stream(req).await + let request_url = req.url.clone(); + let request_headers = req.headers.clone(); + let response = transport.stream(req).await; + observe_auth_response_headers( + auth.as_ref(), + &request_url, + &request_headers, + &response, + ); + response } }, ) @@ -152,3 +170,43 @@ impl EndpointSession { Ok(stream) } } + +fn observe_auth_response_headers( + auth: &dyn crate::auth::AuthProvider, + request_url: &str, + request_headers: &HeaderMap, + response: &Result, +) where + T: ResponseHeaders, +{ + match response { + Ok(response) => { + auth.observe_response_headers(request_url, request_headers, response.headers()) + } + Err(TransportError::Http { + headers: Some(headers), + .. + }) => auth.observe_response_headers(request_url, request_headers, headers), + Err(_) => {} + } +} + +trait ResponseHeaders { + fn headers(&self) -> &HeaderMap; +} + +impl ResponseHeaders for Response { + fn headers(&self) -> &HeaderMap { + &self.headers + } +} + +impl ResponseHeaders for StreamResponse { + fn headers(&self) -> &HeaderMap { + &self.headers + } +} + +#[cfg(test)] +#[path = "session_tests.rs"] +mod tests; diff --git a/codex-rs/codex-api/src/endpoint/session_tests.rs b/codex-rs/codex-api/src/endpoint/session_tests.rs new file mode 100644 index 0000000000..82cf9c0748 --- /dev/null +++ b/codex-rs/codex-api/src/endpoint/session_tests.rs @@ -0,0 +1,56 @@ +use super::*; +use std::sync::Mutex as StdMutex; + +#[derive(Default)] +struct RecordingAuthProvider { + observed_headers: StdMutex>, +} + +impl crate::auth::AuthProvider for RecordingAuthProvider { + fn add_auth_headers(&self, _headers: &mut HeaderMap) {} + + fn observe_response_headers( + &self, + _request_url: &str, + request_headers: &HeaderMap, + response_headers: &HeaderMap, + ) { + self.observed_headers + .lock() + .expect("recording auth lock should not be poisoned") + .push((request_headers.clone(), response_headers.clone())); + } +} + +#[test] +fn observe_auth_response_headers_retains_request_and_http_error_headers() { + let mut request_headers = HeaderMap::new(); + request_headers.insert("x-oai-is", "ois1.sent.nonce.ciphertext".parse().unwrap()); + let mut response_headers = HeaderMap::new(); + response_headers.insert( + "x-oai-is-update", + "ois1.rotated.nonce.ciphertext".parse().unwrap(), + ); + let response: Result = Err(TransportError::Http { + status: http::StatusCode::UNAUTHORIZED, + url: Some("https://chatgpt.com/backend-api/codex/responses".to_string()), + headers: Some(response_headers.clone()), + body: None, + }); + let auth = RecordingAuthProvider::default(); + + observe_auth_response_headers( + &auth, + "https://chatgpt.com/backend-api/codex/responses", + &request_headers, + &response, + ); + + assert_eq!( + auth.observed_headers + .lock() + .expect("recording auth lock should not be poisoned") + .as_slice(), + &[(request_headers, response_headers)] + ); +}