[codex-api] add URL-scoped auth state hooks [ci changed_files]

This commit is contained in:
Cooper Gamble
2026-06-02 19:55:19 +00:00
parent 4019ae532e
commit b8da8fae13
4 changed files with 206 additions and 14 deletions

View File

@@ -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<Request, AuthError> {
let mut request = request;
self.add_auth_headers(&mut request.headers);
self.add_auth_headers_for_url(&request.url, &mut request.headers);
Ok(request)
}
}

View File

@@ -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<ResponsesWebsocketClose> {
let Message::Close(frame) = message else {
return None;
@@ -472,7 +521,17 @@ async fn connect_websocket(
url: Url,
headers: HeaderMap,
turn_state: Option<Arc<OnceLock<String>>>,
) -> Result<(WsStream, StatusCode, bool, Option<String>, Option<String>), ApiError> {
) -> Result<
(
WsStream,
StatusCode,
bool,
Option<String>,
Option<String>,
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(),
))
}

View File

@@ -102,7 +102,16 @@ impl<T: HttpTransport> EndpointSession<T> {
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<T: HttpTransport> EndpointSession<T> {
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<T: HttpTransport> EndpointSession<T> {
Ok(stream)
}
}
fn observe_auth_response_headers<T>(
auth: &dyn crate::auth::AuthProvider,
request_url: &str,
request_headers: &HeaderMap,
response: &Result<T, TransportError>,
) 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;

View File

@@ -0,0 +1,56 @@
use super::*;
use std::sync::Mutex as StdMutex;
#[derive(Default)]
struct RecordingAuthProvider {
observed_headers: StdMutex<Vec<(HeaderMap, HeaderMap)>>,
}
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<Response, TransportError> = 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)]
);
}