use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
struct MockTransport {
stream: String,
requests: Mutex<Vec<HttpRequest>>,
expire_first: bool,
}
impl HttpTransport for MockTransport {
fn stream_json(
&self,
request: HttpRequest,
on_chunk: &mut dyn FnMut(&str) -> anyhow::Result<()>,
) -> anyhow::Result<()> {
self.stream_json_cancellable(request, &AgentCancellation::default(), on_chunk)
}
fn stream_json_cancellable(
&self,
request: HttpRequest,
cancellation: &AgentCancellation,
on_chunk: &mut dyn FnMut(&str) -> anyhow::Result<()>,
) -> anyhow::Result<()> {
self.stream_json_cancellable_with_semantic_deadline(
request,
cancellation,
&AtomicU64::new(0),
on_chunk,
)
}
fn stream_json_cancellable_with_semantic_deadline(
&self,
request: HttpRequest,
cancellation: &AgentCancellation,
_: &AtomicU64,
on_chunk: &mut dyn FnMut(&str) -> anyhow::Result<()>,
) -> anyhow::Result<()> {
let mut requests = self.requests.lock().unwrap();
requests.push(request);
if self.expire_first && requests.len() == 1 {
return Err(crate::providers::error::ProviderError::http_status(401,
format!("provider request failed for {CODEX_RESPONSES_URL}: Your ChatGPT session expired before this request finished.")
).into());
}
drop(requests);
for chunk in self.stream.as_bytes().chunks(17) {
cancellation.check()?;
on_chunk(std::str::from_utf8(chunk).unwrap())?;
}
Ok(())
}
}
fn provider(stream: String) -> OpenAiCodexProvider<MockTransport> {
OpenAiCodexProvider::new(
"active-model",
"test-token",
Some("test-account".into()),
MockTransport {
stream,
requests: Mutex::new(Vec::new()),
expire_first: false,
},
)
}
fn event(value: Value) -> String {
format!("data: {value}\r\n\r\n")
}
fn message() -> Value {
json!({"type":"message", "id":"message-1", "role":"assistant", "content":[{"type":"output_text", "text":"Grounded answer", "annotations":[{"type":"url_citation", "url":"https://example.com/source", "title":"Source", "start_index":0, "end_index":8}]}]})
}
#[test]
fn codex_web_search_uses_active_model_native_tool_and_preserves_citations() {
for final_output in [false, true] {
let item = message();
let mut stream =
event(json!({"type":"response.output_text.delta", "delta":"Grounded answer"}));
if !final_output {
stream += &event(json!({"type":"response.output_item.done", "item":item}));
}
stream += &event(
json!({"type":"response.completed", "response":{"status":"completed", "output":if final_output {vec![item]} else {vec![]}}}),
);
let provider = provider(stream);
let output = provider
.web_search("question", 3, &AgentCancellation::default())
.unwrap();
assert_eq!(output["synthesis"], "Grounded answer");
assert_eq!(output["citations"][0]["url"], "https://example.com/source");
assert_eq!(output["citations"][0]["start_index"], 0);
assert_eq!(output["kind"], "search_synthesis");
let requests = provider.transport.requests.lock().unwrap();
assert_eq!(requests.len(), 1);
let request = &requests[0];
assert_eq!(request.url, CODEX_RESPONSES_URL);
assert_eq!(request.method, "POST");
assert_eq!(request.body["model"], "active-model");
assert_eq!(
request.body["tools"],
json!([{"type":"web_search","external_web_access":true,"search_context_size":"medium"}])
);
assert_eq!(request.body["tool_choice"], "required");
assert_eq!(request.body["store"], false);
assert_eq!(request.body["stream"], true);
assert_eq!(request.body["include"], json!([]));
assert!(request.body["input"].to_string().contains("question"));
assert!(
request
.headers
.iter()
.any(|(key, value)| key.eq_ignore_ascii_case("authorization")
&& value == "Bearer test-token")
);
}
}
#[test]
fn codex_web_search_rejects_missing_completion_failures_and_oversized_streams() {
let delta = event(json!({"type":"response.output_text.delta", "delta":"partial"}));
for ending in [
"data: [DONE]\n\n".to_owned(),
event(json!({"type":"response.failed"})),
event(json!({"type":"response.incomplete"})),
event(json!({"type":"error","message":"failure"})),
] {
assert!(
provider(format!("{delta}{ending}"))
.web_search("question", 5, &AgentCancellation::default())
.is_err()
);
}
let oversized = format!("{delta}:{}\n\n", "x".repeat(MAX_SEARCH_RESPONSE_BYTES));
assert!(
provider(oversized)
.web_search("question", 5, &AgentCancellation::default())
.is_err()
);
}
#[test]
fn codex_web_search_cancellation_prevents_http() {
let provider = provider(String::new());
let (cancellation, handle) = AgentCancellation::default().child_token();
handle.cancel();
assert!(provider.web_search("question", 5, &cancellation).is_err());
assert!(provider.transport.requests.lock().unwrap().is_empty());
}
#[test]
fn codex_web_search_refreshes_expired_session_without_changing_request() {
let stream = event(
json!({"type":"response.completed", "response":{"status":"completed","output":[message()]}}),
);
let mut provider = provider(stream);
provider.transport.expire_first = true;
let refreshes = Arc::new(AtomicUsize::new(0));
let count = refreshes.clone();
let provider = provider.with_auth_refresh(move |_| {
count.fetch_add(1, Ordering::SeqCst);
Ok(CodexRefreshedAuth {
access_token: "refreshed-token".into(),
account_id: Some("test-account".into()),
})
});
provider
.web_search("question", 5, &AgentCancellation::default())
.unwrap();
assert_eq!(refreshes.load(Ordering::SeqCst), 1);
let requests = provider.transport.requests.lock().unwrap();
assert_eq!(requests.len(), 2);
assert_eq!(requests[0].body, requests[1].body);
assert!(
requests[1]
.headers
.iter()
.any(|(key, value)| key.eq_ignore_ascii_case("authorization")
&& value == "Bearer refreshed-token")
);
}
#[test]
fn codex_web_search_parses_many_events_with_all_newline_styles() {
for boundary in ["\n\n", "\r\n\r\n", "\n\r\n", "\r\n\n"] {
let delta =
format!("data: {{\"type\":\"response.output_text.delta\",\"delta\":\"é\"}}{boundary}");
let mut stream = delta.repeat(32_000);
stream += &format!("data: {{\"type\":\"response.completed\"}}{boundary}");
assert!(stream.len() < MAX_SEARCH_RESPONSE_BYTES);
let result = parse_search_response(&stream, || Ok(())).unwrap();
assert_eq!(result["synthesis"], "é".repeat(32_000));
}
}
#[test]
fn codex_web_search_cancels_during_event_processing_and_long_frame_scanning() {
let delta = event(json!({"type":"response.output_text.delta", "delta":"answer"}));
for stream in [
format!("{}data: invalid JSON\n\n", delta.repeat(32_000)),
format!(":{}\n\ndata: invalid JSON\n\n", "x".repeat(100_000)),
] {
let (cancellation, handle) = AgentCancellation::default().child_token();
let mut checks = 0;
let error = parse_search_response(&stream, || {
checks += 1;
if checks == 10 {
handle.cancel();
}
cancellation.check()
})
.unwrap_err();
assert!(crate::cancellation::is_run_canceled(&error), "{error}");
assert_eq!(checks, 10);
}
}