Skip to main content

harness_web/
provider.rs

1//! The search provider a Host supplies: the [`SearchProvider`] trait, the
2//! validated [`SearchQuery`] the search tool hands it, the
3//! [`SearchResults`] it answers with, and the [`SearchError`] it fails
4//! with.
5
6/// A backend that runs web searches for the `promptforge/web/search` tool.
7///
8/// The Host registers one under the key
9/// [`SEARCH_PROVIDER`](crate::SEARCH_PROVIDER). The tool validates the
10/// model's arguments before it calls the provider. The tool leaves the
11/// deadline to the provider, which must limit how long each search takes.
12#[async_trait::async_trait]
13pub trait SearchProvider: Send + Sync {
14    /// Runs `query` and returns its results.
15    ///
16    /// # Errors
17    /// Returns a [`SearchError`] whose kind says whether the transport or
18    /// the backend failed. The model sees the error's message after a
19    /// `web_search: ` prefix, so the message must not contain a credential.
20    async fn search(&self, query: SearchQuery) -> Result<SearchResults, SearchError>;
21}
22
23/// The validated arguments of one `promptforge/web/search` call.
24///
25/// The search tool builds one only from arguments that pass its checks:
26///
27/// - `query` holds more than whitespace and has at most 400 characters.
28/// - `count`, when given, is in `1..=20`.
29/// - `country` and `search_lang`, when given, have 1 to 128 characters.
30/// - Each domain list holds at most 20 hostnames.
31///
32/// An empty domain list keeps every result.
33#[derive(Clone, Debug, Default, PartialEq, Eq)]
34pub struct SearchQuery {
35    /// The search query.
36    pub query: String,
37    /// The number of results wanted.
38    pub count: Option<u8>,
39    /// The freshness filter.
40    pub freshness: Option<Freshness>,
41    /// The country code for the search.
42    pub country: Option<String>,
43    /// The search language code.
44    pub search_lang: Option<String>,
45    /// The SafeSearch level.
46    pub safesearch: Option<SafeSearch>,
47    /// Keep only results from these hostnames.
48    pub include_domains: Vec<String>,
49    /// Drop results from these hostnames.
50    pub exclude_domains: Vec<String>,
51}
52
53/// How recent the results of a [`SearchQuery`] must be.
54///
55/// The search tool reads it from the model's arguments. It refuses any
56/// token other than `pd`, `pw`, `pm`, or `py` as an invalid argument, so
57/// that token never reaches the provider.
58#[derive(Clone, Copy, Debug, PartialEq, Eq, serde::Deserialize)]
59#[serde(rename_all = "lowercase")]
60#[non_exhaustive]
61pub enum Freshness {
62    /// Past day.
63    Pd,
64    /// Past week.
65    Pw,
66    /// Past month.
67    Pm,
68    /// Past year.
69    Py,
70}
71
72impl Freshness {
73    /// Returns the filter's token: `pd`, `pw`, `pm`, or `py`.
74    #[must_use]
75    pub fn as_str(self) -> &'static str {
76        match self {
77            Freshness::Pd => "pd",
78            Freshness::Pw => "pw",
79            Freshness::Pm => "pm",
80            Freshness::Py => "py",
81        }
82    }
83}
84
85/// The SafeSearch filtering level of a [`SearchQuery`].
86///
87/// The search tool reads it from the model's arguments. It refuses any
88/// token other than `off`, `moderate`, or `strict` as an invalid argument,
89/// so that token never reaches the provider.
90#[derive(Clone, Copy, Debug, PartialEq, Eq, serde::Deserialize)]
91#[serde(rename_all = "lowercase")]
92#[non_exhaustive]
93pub enum SafeSearch {
94    /// Filtering turned off.
95    Off,
96    /// Moderate filtering.
97    Moderate,
98    /// Strict filtering.
99    Strict,
100}
101
102impl SafeSearch {
103    /// Returns the level's token: `off`, `moderate`, or `strict`.
104    #[must_use]
105    pub fn as_str(self) -> &'static str {
106        match self {
107            SafeSearch::Off => "off",
108            SafeSearch::Moderate => "moderate",
109            SafeSearch::Strict => "strict",
110        }
111    }
112}
113
114/// The results a [`SearchProvider`] returns for one search, with the query
115/// it ran.
116///
117/// The search tool returns it to the model as compact JSON: an object with
118/// `query` and a `results` array. The fields appear in the same order as
119/// in the Gateway's search output. Like the Gateway, the JSON includes a
120/// result's `age` and `site_name` only when present, and its
121/// `extra_snippets` only when it holds at least one snippet.
122#[derive(Clone, Debug, Default, PartialEq, Eq, serde::Serialize)]
123pub struct SearchResults {
124    /// The query the provider ran.
125    pub query: String,
126    /// The results, in the provider's order.
127    pub results: Vec<SearchResult>,
128}
129
130/// One result in [`SearchResults`]: a title, a URL, a description, and
131/// optional extras.
132///
133/// The search tool rejects the provider's whole reply when any result has
134/// a blank `url`.
135#[derive(Clone, Debug, Default, PartialEq, Eq, serde::Serialize)]
136pub struct SearchResult {
137    /// The result's title.
138    pub title: String,
139    /// The result's URL.
140    pub url: String,
141    /// A short description or snippet.
142    pub description: String,
143    /// The result's age, when the provider reports one.
144    #[serde(skip_serializing_if = "Option::is_none")]
145    pub age: Option<String>,
146    /// The hostname of `url`, when the provider derives one.
147    #[serde(skip_serializing_if = "Option::is_none")]
148    pub site_name: Option<String>,
149    /// Extra snippets from the provider.
150    #[serde(skip_serializing_if = "Vec::is_empty")]
151    pub extra_snippets: Vec<String>,
152}
153
154/// Which part of a search failed: the transport or the backend.
155#[derive(Clone, Copy, Debug, PartialEq, Eq)]
156#[non_exhaustive]
157pub enum SearchErrorKind {
158    /// Sending the request or reading its reply failed. A timeout counts as
159    /// this kind.
160    Transport,
161    /// The backend answered with a failure or with a reply the provider
162    /// rejects.
163    Backend,
164}
165
166/// The error a [`SearchProvider`] returns when a search fails.
167///
168/// It holds a kind, a message, and an optional cause. The search tool
169/// keeps the whole error as the source of the tool error it returns, so
170/// the cause stays in the error chain.
171#[derive(Debug, thiserror::Error)]
172#[error("{message}")]
173pub struct SearchError {
174    kind: SearchErrorKind,
175    message: String,
176    #[source]
177    source: Option<Box<dyn std::error::Error + Send + Sync>>,
178}
179
180impl SearchError {
181    /// Builds an error from a kind and a message alone.
182    #[must_use]
183    pub fn new(kind: SearchErrorKind, message: impl Into<String>) -> SearchError {
184        SearchError {
185            kind,
186            message: message.into(),
187            source: None,
188        }
189    }
190
191    /// Builds an error whose cause is `source`.
192    #[must_use]
193    pub fn with_source(
194        kind: SearchErrorKind,
195        message: impl Into<String>,
196        source: impl std::error::Error + Send + Sync + 'static,
197    ) -> SearchError {
198        SearchError {
199            kind,
200            message: message.into(),
201            source: Some(Box::new(source)),
202        }
203    }
204
205    /// Returns which part of the search failed.
206    #[must_use]
207    pub fn kind(&self) -> SearchErrorKind {
208        self.kind
209    }
210}