mirror of
https://github.com/openai/codex.git
synced 2026-09-06 15:29:32 +00:00
[codex-api] add URL-scoped auth state hooks [ci changed_files]
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
56
codex-rs/codex-api/src/endpoint/session_tests.rs
Normal file
56
codex-rs/codex-api/src/endpoint/session_tests.rs
Normal 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)]
|
||||
);
|
||||
}
|
||||
Reference in New Issue
Block a user