Skip to main content

codex_api/endpoint/
search.rs

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}