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}