Skip to main content

web_search/
search.rs

1//! Web Search Engine - main entry point
2
3use 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/// Status of one provider in a detailed search.
37#[derive(Debug, Clone, Serialize, PartialEq, Eq)]
38#[serde(rename_all = "lowercase")]
39pub enum ProviderOutcomeStatus {
40    /// Provider returned a valid result list (which may be empty).
41    Success,
42    /// Provider returned an error.
43    Error,
44    /// Provider is registered but disabled.
45    Unavailable,
46}
47
48/// Serializable provider error that does not erase its broad category.
49#[derive(Debug, Clone, Serialize)]
50pub struct ProviderError {
51    /// Stable broad error category.
52    pub kind: String,
53    /// Human-readable error detail.
54    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/// Results, diagnostics, and exact captures for one provider.
77#[derive(Debug, Clone, Serialize)]
78#[serde(rename_all = "camelCase")]
79pub struct ProviderOutcome {
80    /// Requested provider id.
81    pub provider: String,
82    /// Success/error/unavailable status.
83    pub status: ProviderOutcomeStatus,
84    /// Unmerged results from this provider.
85    pub results: Vec<SearchResult>,
86    /// Every exact HTTP response observed for this provider.
87    pub responses: Vec<TransportResponse>,
88    /// Structured error, when status is `Error`.
89    #[serde(skip_serializing_if = "Option::is_none")]
90    pub error: Option<ProviderError>,
91}
92
93/// Fused results alongside every provider outcome.
94#[derive(Debug, Clone, Serialize)]
95pub struct DetailedSearchResult {
96    /// Fused successful provider results.
97    pub results: Vec<SearchResult>,
98    /// Per-provider results, errors, and response captures.
99    pub outcomes: Vec<ProviderOutcome>,
100}
101
102/// Configuration for the web search engine
103#[derive(Debug, Clone, Default)]
104pub struct WebSearchConfig {
105    /// Providers to use by default
106    pub providers: Vec<String>,
107    /// Google API key
108    pub google_api_key: Option<String>,
109    /// Google Custom Search Engine ID
110    pub google_cx: Option<String>,
111    /// Bing API key
112    pub bing_api_key: Option<String>,
113    /// Default weights for providers
114    pub weights: HashMap<String, f64>,
115    /// Default merge strategy
116    pub merge_strategy: MergeStrategy,
117}
118
119impl WebSearchConfig {
120    /// Create config from environment variables
121    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
133/// Web Search Engine
134pub 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    /// Create a new web search engine with default configuration
144    pub fn new() -> Self {
145        Self::with_config(WebSearchConfig::from_env())
146    }
147
148    /// Create a new web search engine with custom configuration.
149    ///
150    /// Providers are instantiated from the typed registry (the single source of
151    /// truth), so every catalogued engine — class-based, descriptor-driven, and
152    /// web-capture-backed — is available for selection.
153    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    /// Search across multiple providers
176    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    /// Search with additional merge options
185    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    /// Search through a caller-owned transport and retain per-provider errors
210    /// and exact response bytes. These provider futures are not spawned: when
211    /// the returned future is dropped, all in-flight work is dropped with it.
212    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    /// Search with a single provider
302    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    /// Search one provider through a caller-owned transport.
323    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    /// Get available provider names
344    pub fn get_available_providers(&self) -> Vec<String> {
345        self.providers.keys().cloned().collect()
346    }
347
348    /// Get the full provider registry (metadata for every known provider).
349    pub fn get_registry(&self) -> &[RegistryEntry] {
350        &self.registry
351    }
352
353    /// Get provider status, enriched with registry metadata (category, label,
354    /// CORS readability, access mechanism) so callers see the same shape the
355    /// JavaScript implementation exposes.
356    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    /// Set provider weight
379    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    /// Enable or disable a provider
390    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/// Provider status information, enriched with registry metadata.
408#[derive(Debug, Clone, serde::Serialize)]
409#[serde(rename_all = "camelCase")]
410pub struct ProviderStatus {
411    /// Whether the provider is enabled
412    pub enabled: bool,
413    /// Provider weight for reranking
414    pub weight: f64,
415    /// Provider category (one of the registry categories)
416    #[serde(skip_serializing_if = "Option::is_none")]
417    pub category: Option<String>,
418    /// Human-readable label
419    #[serde(skip_serializing_if = "Option::is_none")]
420    pub label: Option<String>,
421    /// Whether the endpoint is browser-CORS readable
422    #[serde(skip_serializing_if = "Option::is_none")]
423    pub cors_readable: Option<bool>,
424    /// How results are obtained (api, html, hybrid, component, ...)
425    #[serde(skip_serializing_if = "Option::is_none")]
426    pub access: Option<String>,
427}