Skip to main content

harness_gateway_client/
search.rs

1//! The Gateway web search provider: each search is one
2//! `POST {api_root}/tools/web_search` with the bearer key under a fixed
3//! deadline, a bounded and sanitized error body, and a capped success body
4//! parsed into private wire types that mirror the Gateway's request and
5//! response, then mapped into the provider's results and errors.
6
7use 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
16/// The largest error body kept for diagnostics, in bytes.
17const MAX_ERROR_BODY: usize = 2000;
18
19/// The largest successful response body accepted from the gateway, in bytes.
20///
21/// Search results include third-party web content, so the body is bounded
22/// to keep a hostile or misbehaving upstream from returning an unbounded
23/// payload. A body past this cap is rejected rather than silently
24/// truncated, since a truncated JSON document is not a valid result set.
25const MAX_RESPONSE_BODY: usize = 256 * 1024;
26
27/// The deadline applied to every outbound request, body included, so a
28/// stalled gateway cannot hang a search (and thus a run) indefinitely.
29const REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
30
31/// The body of a Gateway web search, mirroring the Gateway's
32/// `POST /v1/tools/web_search` request.
33///
34/// Only `query` is required. An absent option or an empty domain list is
35/// left out of the body, and the Gateway applies its own default.
36#[derive(Clone, Debug, Default, PartialEq, Eq, serde::Serialize)]
37struct GatewaySearchRequest {
38    /// The search query.
39    query: String,
40    /// The number of results wanted; the Gateway clamps it to its maximum.
41    #[serde(skip_serializing_if = "Option::is_none")]
42    count: Option<u8>,
43    /// The freshness filter: `pd`, `pw`, `pm`, `py`, or a
44    /// `YYYY-MM-DDtoYYYY-MM-DD` range.
45    #[serde(skip_serializing_if = "Option::is_none")]
46    freshness: Option<String>,
47    /// The country code for the search.
48    #[serde(skip_serializing_if = "Option::is_none")]
49    country: Option<String>,
50    /// The search language code.
51    #[serde(skip_serializing_if = "Option::is_none")]
52    search_lang: Option<String>,
53    /// The SafeSearch level: `off`, `moderate`, or `strict`.
54    #[serde(skip_serializing_if = "Option::is_none")]
55    safesearch: Option<String>,
56    /// Keep only results from these hostnames.
57    #[serde(skip_serializing_if = "Vec::is_empty")]
58    include_domains: Vec<String>,
59    /// Drop results from these hostnames.
60    #[serde(skip_serializing_if = "Vec::is_empty")]
61    exclude_domains: Vec<String>,
62}
63
64/// The reply to a Gateway web search, mirroring the Gateway's
65/// `POST /v1/tools/web_search` response.
66///
67/// `results` and each result's `url` are required, so a reply without them
68/// is a malformed response. Every other field defaults when absent, and
69/// unknown fields are ignored so the Gateway can grow its reply.
70#[derive(Clone, Debug, PartialEq, Eq, serde::Deserialize)]
71struct GatewaySearchResponse {
72    /// The query the Gateway ran, after its trimming.
73    #[serde(default)]
74    query: String,
75    /// The result rows, in the Gateway's order.
76    results: Vec<GatewaySearchResult>,
77}
78
79/// One row of a [`GatewaySearchResponse`].
80///
81/// The `url` is required but may be empty; judging an empty `url` is left
82/// to the caller.
83#[derive(Clone, Debug, PartialEq, Eq, serde::Deserialize)]
84struct GatewaySearchResult {
85    /// The result's title.
86    #[serde(default)]
87    title: String,
88    /// The result's URL.
89    url: String,
90    /// A short description or snippet.
91    #[serde(default)]
92    description: String,
93    /// The result's age, when the provider reports one.
94    age: Option<String>,
95    /// The hostname of `url`, when the Gateway could derive one.
96    site_name: Option<String>,
97    /// Extra snippets from the provider.
98    #[serde(default)]
99    extra_snippets: Vec<String>,
100}
101
102/// Which side of a Gateway web search failed.
103#[derive(Clone, Copy, Debug, PartialEq, Eq)]
104#[non_exhaustive]
105pub enum GatewaySearchErrorKind {
106    /// Sending the request or reading a successful reply failed, timeouts
107    /// included.
108    Transport,
109    /// The Gateway answered with a failure status or an unusable body.
110    Backend,
111}
112
113/// An error from a failed Gateway web search, with its kind, a message, and
114/// the underlying cause when there is one.
115///
116/// The message starts with what failed, such as `request failed` or
117/// `backend returned 502: ...`. When the message includes the Gateway's
118/// error body, that body is truncated to a size limit and its control
119/// characters are escaped. The message never contains the bearer key.
120#[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    /// Returns which side of the search failed.
151    #[must_use]
152    pub fn kind(&self) -> GatewaySearchErrorKind {
153        self.kind
154    }
155}
156
157/// The [`SearchProvider`] a Host supplies to search the web through the
158/// Gateway.
159///
160/// It sends every search to one Gateway API root with the Gateway's shared
161/// bearer key. Each search is one `POST {api_root}/tools/web_search`
162/// request under a 30-second deadline. The deadline covers the whole
163/// exchange, including reading the reply body. The search vendor's
164/// credential stays in the Gateway. This provider sends only the Gateway's
165/// key.
166#[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    /// Builds a search provider from a validated [`GatewayEndpoint`] and a
186    /// redacted [`SecretString`] bearer key.
187    #[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    /// Runs one web search and returns the Gateway's parsed reply, failing
206    /// as the [`SearchProvider`] impl documents.
207    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            // The error body is external gateway content: bound the read and
231            // sanitize control characters. If the body itself cannot be read,
232            // keep the read failure as the returned error's `source()`.
233            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        // Success bodies hold third-party content: bound them (rejecting cap
256        // overflow), then parse the promised JSON shape.
257        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    /// Runs `query` as one Gateway web search and returns the Gateway's
271    /// results with every field copied as the Gateway sent it.
272    ///
273    /// # Errors
274    /// Returns a [`SearchError`] with the message of the
275    /// [`GatewaySearchError`] it keeps as its source. Its kind is:
276    /// - `Transport` when sending the request fails (`request failed`) or
277    ///   reading a successful reply fails (`reading response failed`),
278    ///   including when the deadline passes;
279    /// - `Backend` when the Gateway answers a failure status
280    ///   (`backend returned {code}: {body}`, or `backend returned {code},
281    ///   and its error body could not be read`), or a success body that is
282    ///   over 256 KiB (`response body exceeded {limit} bytes`), is invalid
283    ///   UTF-8 (`response body was not valid UTF-8`), or fails to parse as a
284    ///   search reply (`malformed search response`).
285    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
292/// The Gateway's request for a validated query.
293fn 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
308/// The provider's results for the Gateway's reply.
309fn 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
327/// The provider's error for a failed Gateway search: its kind and text,
328/// with the Gateway's error as the cause.
329fn 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
337/// Escapes control characters in an external diagnostic body so a hostile
338/// gateway cannot inject terminal/log control sequences or forge multiline
339/// records through an error `Display`.
340fn 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
357/// Reads at most `limit` bytes of a diagnostic body, stopping early once the cap
358/// is reached. Used for the error path, where a truncated, lossy rendering is an
359/// acceptable diagnostic.
360async 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
383/// Reads a success body, rejecting it once it would exceed `limit` bytes rather
384/// than truncating (a truncated JSON document is not a valid result set), and
385/// requiring valid UTF-8.
386async 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;