1use crate::auth::SharedAuthProvider;
2use crate::endpoint::session::EndpointSession;
3use crate::error::ApiError;
4use crate::provider::Provider;
5use crate::search::SearchRequest;
6use crate::search::SearchResponse;
7use codex_client::HttpTransport;
8use codex_client::RequestTelemetry;
9use http::HeaderMap;
10use http::Method;
11use serde_json::to_value;
12use std::sync::Arc;
13
14pub struct SearchClient<T: HttpTransport> {
15 session: EndpointSession<T>,
16}
17
18impl<T: HttpTransport> SearchClient<T> {
19 pub fn new(transport: T, provider: Provider, auth: SharedAuthProvider) -> Self {
20 Self {
21 session: EndpointSession::new(transport, provider, auth),
22 }
23 }
24
25 pub fn with_telemetry(self, request: Option<Arc<dyn RequestTelemetry>>) -> Self {
26 Self {
27 session: self.session.with_request_telemetry(request),
28 }
29 }
30
31 fn path() -> &'static str {
32 "alpha/search"
33 }
34
35 pub async fn search(
36 &self,
37 request: &SearchRequest,
38 extra_headers: HeaderMap,
39 ) -> Result<SearchResponse, ApiError> {
40 let body = to_value(request)
41 .map_err(|e| ApiError::Stream(format!("failed to encode search request: {e}")))?;
42 let resp = self
43 .session
44 .execute(Method::POST, Self::path(), extra_headers, Some(body))
45 .await?;
46 serde_json::from_slice(&resp.body)
47 .map_err(|e| ApiError::Stream(format!("failed to decode search response: {e}")))
48 }
49}
50
51#[cfg(test)]
52mod tests {
53 use super::*;
54 use crate::auth::AuthProvider;
55 use crate::provider::RetryConfig;
56 use crate::search::AllowedCaller;
57 use crate::search::ApproximateLocation;
58 use crate::search::ExternalWebAccess;
59 use crate::search::LocationType;
60 use crate::search::OpenOperation;
61 use crate::search::SearchCommands;
62 use crate::search::SearchContextSize;
63 use crate::search::SearchFilters;
64 use crate::search::SearchImageSettings;
65 use crate::search::SearchInput;
66 use crate::search::SearchQuery;
67 use crate::search::SearchSettings;
68 use codex_client::Request;
69 use codex_client::RequestBody;
70 use codex_client::Response;
71 use codex_client::StreamResponse;
72 use codex_client::TransportError;
73 use codex_protocol::ResponseItemId;
74 use codex_protocol::models::ContentItem;
75 use codex_protocol::models::ResponseItem;
76 use http::StatusCode;
77 use pretty_assertions::assert_eq;
78 use serde_json::json;
79 use std::sync::Mutex;
80 use std::time::Duration;
81
82 #[derive(Clone, Default)]
83 struct DummyAuth;
84
85 impl AuthProvider for DummyAuth {
86 fn add_auth_headers(&self, _headers: &mut HeaderMap) {}
87 }
88
89 #[derive(Clone)]
90 struct CapturingTransport {
91 last_request: Arc<Mutex<Option<Request>>>,
92 response_body: Arc<Vec<u8>>,
93 }
94
95 impl CapturingTransport {
96 fn new(response_body: Vec<u8>) -> Self {
97 Self {
98 last_request: Arc::new(Mutex::new(None)),
99 response_body: Arc::new(response_body),
100 }
101 }
102 }
103
104 impl HttpTransport for CapturingTransport {
105 async fn execute(&self, req: Request) -> Result<Response, TransportError> {
106 *self.last_request.lock().expect("lock request store") = Some(req);
107 Ok(Response {
108 status: StatusCode::OK,
109 headers: HeaderMap::new(),
110 body: self.response_body.as_ref().clone().into(),
111 })
112 }
113
114 async fn stream(&self, _req: Request) -> Result<StreamResponse, TransportError> {
115 Err(TransportError::Build("stream should not run".to_string()))
116 }
117 }
118
119 fn provider() -> Provider {
120 Provider {
121 name: "test".to_string(),
122 base_url: "https://example.com/v1".to_string(),
123 query_params: None,
124 headers: HeaderMap::new(),
125 retry: RetryConfig {
126 max_attempts: 1,
127 base_delay: Duration::from_millis(1),
128 retry_429: false,
129 retry_5xx: true,
130 retry_transport: true,
131 },
132 stream_idle_timeout: Duration::from_secs(1),
133 }
134 }
135
136 #[tokio::test]
137 async fn search_posts_typed_request_and_parses_output() {
138 let transport = CapturingTransport::new(
139 serde_json::to_vec(&json!({
140 "encrypted_output": "ciphertext",
141 "output": "search result",
142 "results": [{
143 "type": "text_result",
144 "ref_id": "turn0search0",
145 "url": "https://example.com/result",
146 "future_field": {"preserved": true},
147 }],
148 }))
149 .expect("serialize response"),
150 );
151 let client = SearchClient::new(transport.clone(), provider(), Arc::new(DummyAuth));
152
153 let response = client
154 .search(
155 &SearchRequest {
156 id: "search-session".to_string(),
157 model: "gpt-test".to_string(),
158 reasoning: None,
159 input: Some(SearchInput::Items(vec![ResponseItem::Message {
160 id: Some(ResponseItemId::with_suffix("msg", "search")),
161 role: "user".to_string(),
162 content: vec![
163 ContentItem::InputText {
164 text: "find this".to_string(),
165 },
166 ContentItem::InputImage {
167 image_url: "https://example.com/image.png".to_string(),
168 detail: None,
169 },
170 ],
171 phase: None,
172 internal_chat_message_metadata_passthrough: None,
173 }])),
174 commands: Some(SearchCommands {
175 search_query: Some(vec![SearchQuery {
176 q: "OpenAI news".to_string(),
177 recency: Some(7),
178 domains: Some(vec!["openai.com".to_string()]),
179 }]),
180 open: Some(vec![OpenOperation {
181 ref_id: "https://openai.com".to_string(),
182 lineno: Some(12),
183 }]),
184 ..Default::default()
185 }),
186 settings: Some(SearchSettings {
187 user_location: Some(ApproximateLocation {
188 r#type: LocationType::Approximate,
189 country: Some("US".to_string()),
190 region: None,
191 city: Some("San Francisco".to_string()),
192 timezone: None,
193 }),
194 search_context_size: Some(SearchContextSize::Low),
195 filters: Some(SearchFilters {
196 allowed_domains: Some(vec!["openai.com".to_string()]),
197 blocked_domains: Some(vec!["example.com".to_string()]),
198 }),
199 image_settings: Some(SearchImageSettings {
200 max_results: Some(4),
201 caption: Some(true),
202 }),
203 allowed_callers: Some(vec![AllowedCaller::Direct]),
204 external_web_access: Some(ExternalWebAccess::Boolean(true)),
205 }),
206 max_output_tokens: Some(2500),
207 },
208 HeaderMap::new(),
209 )
210 .await
211 .expect("search request should succeed");
212
213 assert_eq!(
214 response,
215 SearchResponse {
216 encrypted_output: Some("ciphertext".to_string()),
217 output: "search result".to_string(),
218 results: Some(vec![json!({
219 "type": "text_result",
220 "ref_id": "turn0search0",
221 "url": "https://example.com/result",
222 "future_field": {"preserved": true},
223 })]),
224 }
225 );
226
227 let request = transport
228 .last_request
229 .lock()
230 .expect("lock request store")
231 .clone()
232 .expect("request should be captured");
233 let body = request
234 .body
235 .as_ref()
236 .and_then(RequestBody::json)
237 .expect("request body should be JSON");
238 assert_eq!(
239 body,
240 &json!({
241 "id": "search-session",
242 "model": "gpt-test",
243 "input": [{
244 "type": "message",
245 "id": "msg_search",
246 "role": "user",
247 "content": [
248 {"type": "input_text", "text": "find this"},
249 {
250 "type": "input_image",
251 "image_url": "https://example.com/image.png"
252 }
253 ]
254 }],
255 "commands": {
256 "search_query": [{
257 "q": "OpenAI news",
258 "recency": 7,
259 "domains": ["openai.com"]
260 }],
261 "open": [{"ref_id": "https://openai.com", "lineno": 12}]
262 },
263 "settings": {
264 "user_location": {
265 "type": "approximate",
266 "country": "US",
267 "city": "San Francisco"
268 },
269 "search_context_size": "low",
270 "filters": {
271 "allowed_domains": ["openai.com"],
272 "blocked_domains": ["example.com"]
273 },
274 "image_settings": {"max_results": 4, "caption": true},
275 "allowed_callers": ["direct"],
276 "external_web_access": true
277 },
278 "max_output_tokens": 2500
279 })
280 );
281 }
282 #[test]
283 fn search_response_defaults_missing_results_for_older_endpoints() {
284 let response: SearchResponse = serde_json::from_value(json!({
285 "encrypted_output": null,
286 "output": "search result",
287 }))
288 .expect("response without results should deserialize");
289
290 assert_eq!(
291 response,
292 SearchResponse {
293 encrypted_output: None,
294 output: "search result".to_string(),
295 results: None,
296 }
297 );
298 }
299
300 #[test]
301 fn search_response_preserves_supported_empty_results() {
302 let response: SearchResponse = serde_json::from_value(json!({
303 "encrypted_output": null,
304 "output": "search result",
305 "results": [],
306 }))
307 .expect("response with empty results should deserialize");
308
309 assert_eq!(
310 response,
311 SearchResponse {
312 encrypted_output: None,
313 output: "search result".to_string(),
314 results: Some(Vec::new()),
315 }
316 );
317 }
318}