1use std::collections::HashMap;
4use std::sync::Arc;
5
6use async_trait::async_trait;
7use futures::future::join_all;
8use serde::Serialize;
9use tokio::sync::RwLock;
10
11use crate::error::SearchError;
12use crate::merger::{merge_results, MergeOptions, MergeStrategy};
13use crate::providers::{
14 build_providers, get_default_provider_ids, get_registry, BuildConfig, RegistryEntry,
15 SearchOptions, SearchProvider, SearchResult,
16};
17use crate::transport::{ReqwestTransport, SearchTransport, TransportRequest, TransportResponse};
18
19struct RecordingTransport {
20 inner: Arc<dyn SearchTransport>,
21 responses: Arc<std::sync::Mutex<Vec<TransportResponse>>>,
22}
23
24#[async_trait]
25impl SearchTransport for RecordingTransport {
26 async fn execute(&self, request: TransportRequest) -> Result<TransportResponse, SearchError> {
27 let response = self.inner.execute(request).await?;
28 self.responses
29 .lock()
30 .expect("response capture mutex poisoned")
31 .push(response.clone());
32 Ok(response)
33 }
34}
35
36#[derive(Debug, Clone, Serialize, PartialEq, Eq)]
38#[serde(rename_all = "lowercase")]
39pub enum ProviderOutcomeStatus {
40 Success,
42 Error,
44 Unavailable,
46}
47
48#[derive(Debug, Clone, Serialize)]
50pub struct ProviderError {
51 pub kind: String,
53 pub message: String,
55}
56
57impl From<&SearchError> for ProviderError {
58 fn from(error: &SearchError) -> Self {
59 let kind = match error {
60 SearchError::RequestError(_) => "request",
61 SearchError::Transport(_) => "transport",
62 SearchError::ParseError(_) | SearchError::JsonError(_) => "parse",
63 SearchError::UrlError(_) => "url",
64 SearchError::UnknownProvider(_) => "unknown_provider",
65 SearchError::ProviderDisabled(_) => "provider_disabled",
66 SearchError::ApiError { .. } => "api",
67 SearchError::ConfigError(_) => "configuration",
68 };
69 Self {
70 kind: kind.to_string(),
71 message: error.to_string(),
72 }
73 }
74}
75
76#[derive(Debug, Clone, Serialize)]
78#[serde(rename_all = "camelCase")]
79pub struct ProviderOutcome {
80 pub provider: String,
82 pub status: ProviderOutcomeStatus,
84 pub results: Vec<SearchResult>,
86 pub responses: Vec<TransportResponse>,
88 #[serde(skip_serializing_if = "Option::is_none")]
90 pub error: Option<ProviderError>,
91}
92
93#[derive(Debug, Clone, Serialize)]
95pub struct DetailedSearchResult {
96 pub results: Vec<SearchResult>,
98 pub outcomes: Vec<ProviderOutcome>,
100}
101
102#[derive(Debug, Clone, Default)]
104pub struct WebSearchConfig {
105 pub providers: Vec<String>,
107 pub google_api_key: Option<String>,
109 pub google_cx: Option<String>,
111 pub bing_api_key: Option<String>,
113 pub weights: HashMap<String, f64>,
115 pub merge_strategy: MergeStrategy,
117}
118
119impl WebSearchConfig {
120 pub fn from_env() -> Self {
122 Self {
123 providers: get_default_provider_ids(),
124 google_api_key: std::env::var("GOOGLE_API_KEY").ok(),
125 google_cx: std::env::var("GOOGLE_CX").ok(),
126 bing_api_key: std::env::var("BING_API_KEY").ok(),
127 weights: HashMap::new(),
128 merge_strategy: MergeStrategy::Rrf,
129 }
130 }
131}
132
133pub struct WebSearchEngine {
135 providers: HashMap<String, Arc<RwLock<Box<dyn SearchProvider>>>>,
136 registry: Vec<RegistryEntry>,
137 default_providers: Vec<String>,
138 default_weights: HashMap<String, f64>,
139 default_strategy: MergeStrategy,
140}
141
142impl WebSearchEngine {
143 pub fn new() -> Self {
145 Self::with_config(WebSearchConfig::from_env())
146 }
147
148 pub fn with_config(config: WebSearchConfig) -> Self {
154 let mut providers: HashMap<String, Arc<RwLock<Box<dyn SearchProvider>>>> = HashMap::new();
155
156 let build_config = BuildConfig {
157 google_api_key: config.google_api_key,
158 google_cx: config.google_cx,
159 bing_api_key: config.bing_api_key,
160 };
161
162 for (id, provider) in build_providers(&build_config) {
163 providers.insert(id, Arc::new(RwLock::new(provider)));
164 }
165
166 Self {
167 providers,
168 registry: get_registry(),
169 default_providers: config.providers,
170 default_weights: config.weights,
171 default_strategy: config.merge_strategy,
172 }
173 }
174
175 pub async fn search(
177 &self,
178 query: &str,
179 options: SearchOptions,
180 ) -> Result<Vec<SearchResult>, SearchError> {
181 self.search_with_options(query, options, None, None).await
182 }
183
184 pub async fn search_with_options(
186 &self,
187 query: &str,
188 options: SearchOptions,
189 providers: Option<Vec<String>>,
190 merge_options: Option<MergeOptions>,
191 ) -> Result<Vec<SearchResult>, SearchError> {
192 let detailed = self
193 .search_detailed_with_options(
194 query,
195 options,
196 providers,
197 merge_options,
198 Arc::new(ReqwestTransport::default()),
199 )
200 .await;
201 for outcome in &detailed.outcomes {
202 if let Some(error) = &outcome.error {
203 tracing::error!("Provider {} failed: {}", outcome.provider, error.message);
204 }
205 }
206 Ok(detailed.results)
207 }
208
209 pub async fn search_detailed_with_options(
213 &self,
214 query: &str,
215 options: SearchOptions,
216 providers: Option<Vec<String>>,
217 merge_options: Option<MergeOptions>,
218 transport: Arc<dyn SearchTransport>,
219 ) -> DetailedSearchResult {
220 if query.is_empty() {
221 return DetailedSearchResult {
222 results: Vec::new(),
223 outcomes: Vec::new(),
224 };
225 }
226
227 let providers_to_use = providers.unwrap_or_else(|| self.default_providers.clone());
228 let merge_opts = merge_options.unwrap_or_else(|| MergeOptions {
229 strategy: self.default_strategy,
230 weights: self.default_weights.clone(),
231 rrf_k: None,
232 remove_duplicates: true,
233 });
234 let futures = providers_to_use.into_iter().map(|name| {
235 let provider = self.providers.get(&name).cloned();
236 let options = options.clone();
237 let transport = transport.clone();
238 async move {
239 let Some(provider) = provider else {
240 let error = SearchError::UnknownProvider(name.clone());
241 return ProviderOutcome {
242 provider: name,
243 status: ProviderOutcomeStatus::Error,
244 results: Vec::new(),
245 responses: Vec::new(),
246 error: Some(ProviderError::from(&error)),
247 };
248 };
249 let provider = provider.read().await;
250 if !provider.is_available() {
251 return ProviderOutcome {
252 provider: name,
253 status: ProviderOutcomeStatus::Unavailable,
254 results: Vec::new(),
255 responses: Vec::new(),
256 error: None,
257 };
258 }
259 let responses = Arc::new(std::sync::Mutex::new(Vec::new()));
260 let recording = RecordingTransport {
261 inner: transport,
262 responses: responses.clone(),
263 };
264 let result = provider
265 .search_with_transport(query, &options, &recording)
266 .await;
267 let captures = responses
268 .lock()
269 .expect("response capture mutex poisoned")
270 .clone();
271 match result {
272 Ok(results) => ProviderOutcome {
273 provider: name,
274 status: ProviderOutcomeStatus::Success,
275 results,
276 responses: captures,
277 error: None,
278 },
279 Err(error) => ProviderOutcome {
280 provider: name,
281 status: ProviderOutcomeStatus::Error,
282 results: Vec::new(),
283 responses: captures,
284 error: Some(ProviderError::from(&error)),
285 },
286 }
287 }
288 });
289 let outcomes = join_all(futures).await;
290 let results_by_provider = outcomes
291 .iter()
292 .filter(|outcome| outcome.status == ProviderOutcomeStatus::Success)
293 .map(|outcome| (outcome.provider.clone(), outcome.results.clone()))
294 .collect();
295 DetailedSearchResult {
296 results: merge_results(&results_by_provider, &merge_opts),
297 outcomes,
298 }
299 }
300
301 pub async fn search_single(
303 &self,
304 query: &str,
305 provider_name: &str,
306 options: SearchOptions,
307 ) -> Result<Vec<SearchResult>, SearchError> {
308 let provider = self
309 .providers
310 .get(provider_name)
311 .ok_or_else(|| SearchError::UnknownProvider(provider_name.to_string()))?;
312
313 let provider = provider.read().await;
314
315 if !provider.is_available() {
316 return Err(SearchError::ProviderDisabled(provider_name.to_string()));
317 }
318
319 provider.search(query, &options).await
320 }
321
322 pub async fn search_single_with_transport(
324 &self,
325 query: &str,
326 provider_name: &str,
327 options: SearchOptions,
328 transport: &dyn SearchTransport,
329 ) -> Result<Vec<SearchResult>, SearchError> {
330 let provider = self
331 .providers
332 .get(provider_name)
333 .ok_or_else(|| SearchError::UnknownProvider(provider_name.to_string()))?;
334 let provider = provider.read().await;
335 if !provider.is_available() {
336 return Err(SearchError::ProviderDisabled(provider_name.to_string()));
337 }
338 provider
339 .search_with_transport(query, &options, transport)
340 .await
341 }
342
343 pub fn get_available_providers(&self) -> Vec<String> {
345 self.providers.keys().cloned().collect()
346 }
347
348 pub fn get_registry(&self) -> &[RegistryEntry] {
350 &self.registry
351 }
352
353 pub async fn get_provider_status(&self) -> HashMap<String, ProviderStatus> {
357 let mut status = HashMap::new();
358
359 for (name, provider) in &self.providers {
360 let p = provider.read().await;
361 let meta = self.registry.iter().find(|e| &e.id == name);
362 status.insert(
363 name.clone(),
364 ProviderStatus {
365 enabled: p.is_available(),
366 weight: p.weight(),
367 category: meta.map(|m| m.category.clone()),
368 label: meta.map(|m| m.label.clone()),
369 cors_readable: meta.map(|m| m.cors_readable),
370 access: meta.map(|m| m.access.clone()),
371 },
372 );
373 }
374
375 status
376 }
377
378 pub async fn set_provider_weight(&self, name: &str, weight: f64) -> Result<(), SearchError> {
380 let provider = self
381 .providers
382 .get(name)
383 .ok_or_else(|| SearchError::UnknownProvider(name.to_string()))?;
384
385 provider.write().await.set_weight(weight);
386 Ok(())
387 }
388
389 pub async fn set_provider_enabled(&self, name: &str, enabled: bool) -> Result<(), SearchError> {
391 let provider = self
392 .providers
393 .get(name)
394 .ok_or_else(|| SearchError::UnknownProvider(name.to_string()))?;
395
396 provider.write().await.set_enabled(enabled);
397 Ok(())
398 }
399}
400
401impl Default for WebSearchEngine {
402 fn default() -> Self {
403 Self::new()
404 }
405}
406
407#[derive(Debug, Clone, serde::Serialize)]
409#[serde(rename_all = "camelCase")]
410pub struct ProviderStatus {
411 pub enabled: bool,
413 pub weight: f64,
415 #[serde(skip_serializing_if = "Option::is_none")]
417 pub category: Option<String>,
418 #[serde(skip_serializing_if = "Option::is_none")]
420 pub label: Option<String>,
421 #[serde(skip_serializing_if = "Option::is_none")]
423 pub cors_readable: Option<bool>,
424 #[serde(skip_serializing_if = "Option::is_none")]
426 pub access: Option<String>,
427}