Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 25 additions & 5 deletions codex-rs/codex-api/src/endpoint/responses.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ use crate::requests::headers::insert_header;
use crate::requests::headers::subagent_header;
use crate::sse::spawn_response_stream;
use crate::telemetry::SseTelemetry;
use codex_client::EncodedJsonBody;
use codex_client::HttpTransport;
use codex_client::RequestCompression;
use codex_client::RequestTelemetry;
Expand Down Expand Up @@ -81,11 +82,16 @@ impl<T: HttpTransport> ResponsesClient<T> {
turn_state,
} = options;

let mut body = serde_json::to_value(&request)
.map_err(|e| ApiError::Stream(format!("failed to encode responses request: {e}")))?;
if request.store && self.session.provider().is_azure_responses_endpoint() {
let body = if request.store && self.session.provider().is_azure_responses_endpoint() {
let mut body = serde_json::to_value(&request).map_err(|e| {
ApiError::Stream(format!("failed to encode responses request: {e}"))
})?;
attach_item_ids(&mut body, &request.input);
EncodedJsonBody::encode(&body)
} else {
EncodedJsonBody::encode(&request)
}
.map_err(|e| ApiError::Stream(format!("failed to encode responses request: {e}")))?;

let mut headers = extra_headers;
if let Some(ref thread_id) = thread_id {
Expand All @@ -96,7 +102,8 @@ impl<T: HttpTransport> ResponsesClient<T> {
insert_header(&mut headers, "x-openai-subagent", &subagent);
}

self.stream(body, headers, compression, turn_state).await
self.stream_encoded(body, headers, compression, turn_state)
.await
}

fn path() -> &'static str {
Expand All @@ -120,6 +127,19 @@ impl<T: HttpTransport> ResponsesClient<T> {
extra_headers: HeaderMap,
compression: Compression,
turn_state: Option<Arc<OnceLock<String>>>,
) -> Result<ResponseStream, ApiError> {
let body = EncodedJsonBody::encode(&body)
.map_err(|e| ApiError::Stream(format!("failed to encode responses request: {e}")))?;
self.stream_encoded(body, extra_headers, compression, turn_state)
.await
}

async fn stream_encoded(
&self,
body: EncodedJsonBody,
extra_headers: HeaderMap,
compression: Compression,
turn_state: Option<Arc<OnceLock<String>>>,
) -> Result<ResponseStream, ApiError> {
let request_compression = match compression {
Compression::None => RequestCompression::None,
Expand All @@ -128,7 +148,7 @@ impl<T: HttpTransport> ResponsesClient<T> {

let stream_response = self
.session
.stream_with(
.stream_encoded_json_with(
Method::POST,
Self::path(),
extra_headers,
Expand Down
22 changes: 12 additions & 10 deletions codex-rs/codex-api/src/endpoint/session.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ use crate::auth::SharedAuthProvider;
use crate::error::ApiError;
use crate::provider::Provider;
use crate::telemetry::run_with_request_telemetry;
use codex_client::EncodedJsonBody;
use codex_client::HttpTransport;
use codex_client::Request;
use codex_client::RequestBody;
Expand Down Expand Up @@ -49,12 +50,12 @@ impl<T: HttpTransport> EndpointSession<T> {
method: &Method,
path: &str,
extra_headers: &HeaderMap,
body: Option<&Value>,
body: Option<&RequestBody>,
) -> Request {
let mut req = self.provider.build_request(method.clone(), path);
req.headers.extend(extra_headers.clone());
if let Some(body) = body {
req.body = Some(RequestBody::Json(body.clone()));
req.body = Some(body.clone());
}
req
}
Expand Down Expand Up @@ -87,6 +88,7 @@ impl<T: HttpTransport> EndpointSession<T> {
where
C: Fn(&mut Request),
{
let body = body.map(RequestBody::Json);
let make_request = || {
let mut req = self.make_request(&method, path, &extra_headers, body.as_ref());
configure(&mut req);
Expand All @@ -112,27 +114,27 @@ impl<T: HttpTransport> EndpointSession<T> {
}

#[instrument(
name = "endpoint_session.stream_with",
name = "endpoint_session.stream_encoded_json_with",
level = "info",
skip_all,
fields(http.method = %method, api.path = path)
)]
pub(crate) async fn stream_with<C>(
pub(crate) async fn stream_encoded_json_with<C>(
&self,
method: Method,
path: &str,
extra_headers: HeaderMap,
body: Option<Value>,
body: Option<EncodedJsonBody>,
configure: C,
) -> Result<StreamResponse, ApiError>
where
C: Fn(&mut Request),
{
let make_request = || {
let mut req = self.make_request(&method, path, &extra_headers, body.as_ref());
configure(&mut req);
req
};
let body = body.map(RequestBody::EncodedJson);
let mut request = self.make_request(&method, path, &extra_headers, body.as_ref());
configure(&mut request);
let request = request.into_prepared().map_err(TransportError::Build)?;
let make_request = || request.clone();

let stream = run_with_request_telemetry(
self.provider.retry.to_policy(),
Expand Down
116 changes: 102 additions & 14 deletions codex-rs/codex-api/tests/clients.rs
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,13 @@ fn assert_path_ends_with(requests: &[Request], suffix: &str) {
);
}

fn request_body_bytes(request: &Request) -> &[u8] {
let Some(RequestBody::EncodedJson(body)) = request.body.as_ref() else {
panic!("expected a prepared request body");
};
body.as_bytes()
}

#[derive(Debug, Default, Clone)]
struct RecordingState {
stream_requests: Arc<Mutex<Vec<Request>>>,
Expand Down Expand Up @@ -138,9 +145,15 @@ fn provider(name: &str) -> Provider {
}
}

#[derive(Debug, Default)]
struct FlakyTransportState {
attempts: i64,
requests: Vec<(RequestBody, HeaderMap, codex_client::RequestCompression)>,
}

#[derive(Clone)]
struct FlakyTransport {
state: Arc<Mutex<i64>>,
state: Arc<Mutex<FlakyTransportState>>,
}

impl Default for FlakyTransport {
Expand All @@ -152,15 +165,23 @@ impl Default for FlakyTransport {
impl FlakyTransport {
fn new() -> Self {
Self {
state: Arc::new(Mutex::new(0)),
state: Arc::new(Mutex::new(FlakyTransportState::default())),
}
}

fn attempts(&self) -> i64 {
*self
.state
self.state
.lock()
.unwrap_or_else(|err| panic!("mutex poisoned: {err}"))
.attempts
}

fn requests(&self) -> Vec<(RequestBody, HeaderMap, codex_client::RequestCompression)> {
self.state
.lock()
.unwrap_or_else(|err| panic!("mutex poisoned: {err}"))
.requests
.clone()
}
}

Expand Down Expand Up @@ -225,14 +246,20 @@ impl HttpTransport for FlakyTransport {
Err(TransportError::Build("execute should not run".to_string()))
}

async fn stream(&self, _req: Request) -> Result<StreamResponse, TransportError> {
let mut attempts = self
async fn stream(&self, req: Request) -> Result<StreamResponse, TransportError> {
let Some(body) = req.body.clone() else {
panic!("request should have a body");
};
let mut state = self
.state
.lock()
.unwrap_or_else(|err| panic!("mutex poisoned: {err}"));
*attempts += 1;
state.attempts += 1;
state
.requests
.push((body, req.headers.clone(), req.compression));

if *attempts == 1 {
if state.attempts == 1 {
return Err(TransportError::Network("first attempt fails".to_string()));
}

Expand Down Expand Up @@ -272,6 +299,51 @@ async fn responses_client_uses_responses_path() -> Result<()> {
Ok(())
}

#[tokio::test]
async fn responses_client_stream_request_preserves_exact_json_body() -> Result<()> {
let state = RecordingState::default();
let transport = RecordingTransport::new(state.clone());
let client = ResponsesClient::new(transport, provider("openai"), Arc::new(NoAuth));
let request = ResponsesApiRequest {
model: "gpt-test".into(),
instructions: "Say hi".into(),
input: vec![ResponseItem::Message {
id: Some("msg_1".into()),
role: "user".into(),
content: vec![ContentItem::InputText { text: "hi".into() }],
phase: None,
}],
tools: Vec::new(),
tool_choice: "auto".into(),
parallel_tool_calls: false,
reasoning: None,
store: false,
stream: true,
include: Vec::new(),
service_tier: None,
prompt_cache_key: None,
text: None,
client_metadata: None,
};
let expected = serde_json::to_vec(&request)?;

let _stream = client
.stream_request(request, ResponsesOptions::default())
.await?;

let requests = state.take_stream_requests();
assert_eq!(requests.len(), 1);
let prepared = requests[0]
.prepare_body_for_send()
.expect("body should prepare");
assert_eq!(prepared.body.as_deref(), Some(expected.as_slice()));
assert_eq!(
prepared.headers.get(http::header::CONTENT_TYPE),
Some(&HeaderValue::from_static("application/json"))
);
Ok(())
}

#[tokio::test]
async fn streaming_client_adds_auth_headers() -> Result<()> {
let state = RecordingState::default();
Expand Down Expand Up @@ -342,12 +414,30 @@ async fn streaming_client_retries_on_transport_error() -> Result<()> {
.stream_request(
request,
ResponsesOptions {
compression: Compression::None,
compression: Compression::Zstd,
..Default::default()
},
)
.await?;
assert_eq!(transport.attempts(), 2);
let requests = transport.requests();
assert_eq!(requests.len(), 2);
assert_eq!(requests[0], requests[1]);
let RequestBody::EncodedJson(first_body) = &requests[0].0 else {
panic!("expected an encoded JSON body");
};
let RequestBody::EncodedJson(second_body) = &requests[1].0 else {
panic!("expected an encoded JSON body");
};
assert_eq!(
first_body.as_bytes().as_ptr(),
second_body.as_bytes().as_ptr()
);
assert_eq!(
requests[0].1.get(http::header::CONTENT_ENCODING),
Some(&HeaderValue::from_static("zstd"))
);
assert_eq!(requests[0].2, codex_client::RequestCompression::None);
Ok(())
}

Expand Down Expand Up @@ -485,11 +575,9 @@ async fn azure_default_store_attaches_ids_and_headers() -> Result<()> {
Some("present")
);

let input_id = req
.body
.as_ref()
.and_then(RequestBody::json)
.and_then(|body| body.get("input"))
let body: serde_json::Value = serde_json::from_slice(request_body_bytes(req))?;
let input_id = body
.get("input")
.and_then(|input| input.get(0))
.and_then(|item| item.get("id"))
.and_then(|id| id.as_str());
Expand Down
1 change: 1 addition & 0 deletions codex-rs/codex-client/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@ pub use crate::default_client::CodexHttpClient;
pub use crate::default_client::CodexRequestBuilder;
pub use crate::error::StreamError;
pub use crate::error::TransportError;
pub use crate::request::EncodedJsonBody;
pub use crate::request::PreparedRequestBody;
pub use crate::request::Request;
pub use crate::request::RequestBody;
Expand Down
Loading
Loading