1use std::fmt;
8use std::time::Duration;
9
10use harness_web::{
11 SearchError, SearchErrorKind, SearchProvider, SearchQuery, SearchResult, SearchResults,
12};
13
14use crate::config::{GatewayEndpoint, SecretString};
15
16const MAX_ERROR_BODY: usize = 2000;
18
19const MAX_RESPONSE_BODY: usize = 256 * 1024;
26
27const REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
30
31#[derive(Clone, Debug, Default, PartialEq, Eq, serde::Serialize)]
37struct GatewaySearchRequest {
38 query: String,
40 #[serde(skip_serializing_if = "Option::is_none")]
42 count: Option<u8>,
43 #[serde(skip_serializing_if = "Option::is_none")]
46 freshness: Option<String>,
47 #[serde(skip_serializing_if = "Option::is_none")]
49 country: Option<String>,
50 #[serde(skip_serializing_if = "Option::is_none")]
52 search_lang: Option<String>,
53 #[serde(skip_serializing_if = "Option::is_none")]
55 safesearch: Option<String>,
56 #[serde(skip_serializing_if = "Vec::is_empty")]
58 include_domains: Vec<String>,
59 #[serde(skip_serializing_if = "Vec::is_empty")]
61 exclude_domains: Vec<String>,
62}
63
64#[derive(Clone, Debug, PartialEq, Eq, serde::Deserialize)]
71struct GatewaySearchResponse {
72 #[serde(default)]
74 query: String,
75 results: Vec<GatewaySearchResult>,
77}
78
79#[derive(Clone, Debug, PartialEq, Eq, serde::Deserialize)]
84struct GatewaySearchResult {
85 #[serde(default)]
87 title: String,
88 url: String,
90 #[serde(default)]
92 description: String,
93 age: Option<String>,
95 site_name: Option<String>,
97 #[serde(default)]
99 extra_snippets: Vec<String>,
100}
101
102#[derive(Clone, Copy, Debug, PartialEq, Eq)]
104#[non_exhaustive]
105pub enum GatewaySearchErrorKind {
106 Transport,
109 Backend,
111}
112
113#[derive(Debug, thiserror::Error)]
121#[error("{message}")]
122pub struct GatewaySearchError {
123 kind: GatewaySearchErrorKind,
124 message: String,
125 #[source]
126 source: Option<Box<dyn std::error::Error + Send + Sync>>,
127}
128
129impl GatewaySearchError {
130 fn new(kind: GatewaySearchErrorKind, message: impl Into<String>) -> GatewaySearchError {
131 GatewaySearchError {
132 kind,
133 message: message.into(),
134 source: None,
135 }
136 }
137
138 fn with_source(
139 kind: GatewaySearchErrorKind,
140 message: impl Into<String>,
141 source: impl std::error::Error + Send + Sync + 'static,
142 ) -> GatewaySearchError {
143 GatewaySearchError {
144 kind,
145 message: message.into(),
146 source: Some(Box::new(source)),
147 }
148 }
149
150 #[must_use]
152 pub fn kind(&self) -> GatewaySearchErrorKind {
153 self.kind
154 }
155}
156
157#[derive(Clone)]
167#[non_exhaustive]
168pub struct GatewaySearch {
169 http: reqwest::Client,
170 base_url: String,
171 key: SecretString,
172 timeout: Duration,
173}
174
175impl fmt::Debug for GatewaySearch {
176 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
177 f.debug_struct("GatewaySearch")
178 .field("base_url", &self.base_url)
179 .field("key", &"<redacted>")
180 .finish_non_exhaustive()
181 }
182}
183
184impl GatewaySearch {
185 #[must_use]
188 pub fn new(endpoint: GatewayEndpoint, key: SecretString) -> GatewaySearch {
189 GatewaySearch::with_timeout(endpoint, key, REQUEST_TIMEOUT)
190 }
191
192 fn with_timeout(
193 endpoint: GatewayEndpoint,
194 key: SecretString,
195 timeout: Duration,
196 ) -> GatewaySearch {
197 GatewaySearch {
198 http: reqwest::Client::new(),
199 base_url: endpoint.url,
200 key,
201 timeout,
202 }
203 }
204
205 async fn search(
208 &self,
209 request: &GatewaySearchRequest,
210 ) -> Result<GatewaySearchResponse, GatewaySearchError> {
211 let response = self
212 .http
213 .post(format!("{}/tools/web_search", self.base_url))
214 .bearer_auth(self.key.expose())
215 .timeout(self.timeout)
216 .json(request)
217 .send()
218 .await
219 .map_err(|source| {
220 GatewaySearchError::with_source(
221 GatewaySearchErrorKind::Transport,
222 "request failed",
223 source,
224 )
225 })?;
226
227 let status = response.status();
228 if !status.is_success() {
229 let code = status.as_u16();
230 match read_bounded(response, MAX_ERROR_BODY).await {
234 Ok(body) => {
235 let body = if body.is_empty() {
236 "(empty body)".to_owned()
237 } else {
238 sanitize_diagnostic(&body)
239 };
240 return Err(GatewaySearchError::new(
241 GatewaySearchErrorKind::Backend,
242 format!("backend returned {code}: {body}"),
243 ));
244 }
245 Err(source) => {
246 return Err(GatewaySearchError::with_source(
247 GatewaySearchErrorKind::Backend,
248 format!("backend returned {code}, and its error body could not be read"),
249 source,
250 ));
251 }
252 }
253 }
254
255 let body = read_capped(response, MAX_RESPONSE_BODY).await?;
258 serde_json::from_str(&body).map_err(|source| {
259 GatewaySearchError::with_source(
260 GatewaySearchErrorKind::Backend,
261 "malformed search response",
262 source,
263 )
264 })
265 }
266}
267
268#[async_trait::async_trait]
269impl SearchProvider for GatewaySearch {
270 async fn search(&self, query: SearchQuery) -> Result<SearchResults, SearchError> {
286 let request = gateway_request(query);
287 let response = self.search(&request).await.map_err(search_error)?;
288 Ok(search_results(response))
289 }
290}
291
292fn gateway_request(query: SearchQuery) -> GatewaySearchRequest {
294 GatewaySearchRequest {
295 query: query.query,
296 count: query.count,
297 freshness: query
298 .freshness
299 .map(|freshness| freshness.as_str().to_owned()),
300 country: query.country,
301 search_lang: query.search_lang,
302 safesearch: query.safesearch.map(|level| level.as_str().to_owned()),
303 include_domains: query.include_domains,
304 exclude_domains: query.exclude_domains,
305 }
306}
307
308fn search_results(response: GatewaySearchResponse) -> SearchResults {
310 SearchResults {
311 query: response.query,
312 results: response
313 .results
314 .into_iter()
315 .map(|result| SearchResult {
316 title: result.title,
317 url: result.url,
318 description: result.description,
319 age: result.age,
320 site_name: result.site_name,
321 extra_snippets: result.extra_snippets,
322 })
323 .collect(),
324 }
325}
326
327fn search_error(error: GatewaySearchError) -> SearchError {
330 let kind = match error.kind {
331 GatewaySearchErrorKind::Backend => SearchErrorKind::Backend,
332 GatewaySearchErrorKind::Transport => SearchErrorKind::Transport,
333 };
334 SearchError::with_source(kind, error.to_string(), error)
335}
336
337fn sanitize_diagnostic(body: &str) -> String {
341 let mut out = String::with_capacity(body.len());
342 for c in body.chars() {
343 match c {
344 '\n' => out.push_str("\\n"),
345 '\r' => out.push_str("\\r"),
346 '\t' => out.push_str("\\t"),
347 c if c.is_control() => {
348 use std::fmt::Write as _;
349 let _ = write!(out, "\\u{{{:04x}}}", u32::from(c));
350 }
351 c => out.push(c),
352 }
353 }
354 out
355}
356
357async fn read_bounded(
361 mut response: reqwest::Response,
362 limit: usize,
363) -> Result<String, GatewaySearchError> {
364 let mut buffer: Vec<u8> = Vec::new();
365 while buffer.len() < limit {
366 let chunk = response.chunk().await.map_err(|source| {
367 GatewaySearchError::with_source(
368 GatewaySearchErrorKind::Transport,
369 "reading response failed",
370 source,
371 )
372 })?;
373 let Some(chunk) = chunk else { break };
374 let take = (limit - buffer.len()).min(chunk.len());
375 buffer.extend_from_slice(&chunk[..take]);
376 if take < chunk.len() {
377 break;
378 }
379 }
380 Ok(String::from_utf8_lossy(&buffer).into_owned())
381}
382
383async fn read_capped(
387 mut response: reqwest::Response,
388 limit: usize,
389) -> Result<String, GatewaySearchError> {
390 let mut buffer: Vec<u8> = Vec::new();
391 while let Some(chunk) = response.chunk().await.map_err(|source| {
392 GatewaySearchError::with_source(
393 GatewaySearchErrorKind::Transport,
394 "reading response failed",
395 source,
396 )
397 })? {
398 if buffer.len() + chunk.len() > limit {
399 return Err(GatewaySearchError::new(
400 GatewaySearchErrorKind::Backend,
401 format!("response body exceeded {limit} bytes"),
402 ));
403 }
404 buffer.extend_from_slice(&chunk);
405 }
406 String::from_utf8(buffer).map_err(|source| {
407 GatewaySearchError::with_source(
408 GatewaySearchErrorKind::Backend,
409 "response body was not valid UTF-8",
410 source,
411 )
412 })
413}
414
415#[cfg(test)]
416#[path = "search-tests.rs"]
417mod tests;