use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use crate::error::SearchError;
use crate::merger::{merge_results, MergeOptions, MergeStrategy};
use crate::providers::{
build_providers, get_default_provider_ids, get_registry, BuildConfig, RegistryEntry,
SearchOptions, SearchProvider, SearchResult,
};
#[derive(Debug, Clone, Default)]
pub struct WebSearchConfig {
pub providers: Vec<String>,
pub google_api_key: Option<String>,
pub google_cx: Option<String>,
pub bing_api_key: Option<String>,
pub weights: HashMap<String, f64>,
pub merge_strategy: MergeStrategy,
}
impl WebSearchConfig {
pub fn from_env() -> Self {
Self {
providers: get_default_provider_ids(),
google_api_key: std::env::var("GOOGLE_API_KEY").ok(),
google_cx: std::env::var("GOOGLE_CX").ok(),
bing_api_key: std::env::var("BING_API_KEY").ok(),
weights: HashMap::new(),
merge_strategy: MergeStrategy::Rrf,
}
}
}
pub struct WebSearchEngine {
providers: HashMap<String, Arc<RwLock<Box<dyn SearchProvider>>>>,
registry: Vec<RegistryEntry>,
default_providers: Vec<String>,
default_weights: HashMap<String, f64>,
default_strategy: MergeStrategy,
}
impl WebSearchEngine {
pub fn new() -> Self {
Self::with_config(WebSearchConfig::from_env())
}
pub fn with_config(config: WebSearchConfig) -> Self {
let mut providers: HashMap<String, Arc<RwLock<Box<dyn SearchProvider>>>> = HashMap::new();
let build_config = BuildConfig {
google_api_key: config.google_api_key,
google_cx: config.google_cx,
bing_api_key: config.bing_api_key,
};
for (id, provider) in build_providers(&build_config) {
providers.insert(id, Arc::new(RwLock::new(provider)));
}
Self {
providers,
registry: get_registry(),
default_providers: config.providers,
default_weights: config.weights,
default_strategy: config.merge_strategy,
}
}
pub async fn search(
&self,
query: &str,
options: SearchOptions,
) -> Result<Vec<SearchResult>, SearchError> {
self.search_with_options(query, options, None, None).await
}
pub async fn search_with_options(
&self,
query: &str,
options: SearchOptions,
providers: Option<Vec<String>>,
merge_options: Option<MergeOptions>,
) -> Result<Vec<SearchResult>, SearchError> {
if query.is_empty() {
return Ok(Vec::new());
}
let providers_to_use = providers.unwrap_or_else(|| self.default_providers.clone());
let merge_opts = merge_options.unwrap_or_else(|| MergeOptions {
strategy: self.default_strategy,
weights: self.default_weights.clone(),
rrf_k: None,
remove_duplicates: true,
});
let mut handles = Vec::new();
for provider_name in &providers_to_use {
let provider = self.providers.get(provider_name).cloned();
if provider.is_none() {
continue;
}
let provider = provider.unwrap();
let query = query.to_string();
let opts = options.clone();
let name = provider_name.clone();
handles.push(tokio::spawn(async move {
let provider = provider.read().await;
if !provider.is_available() {
return (name, Vec::new());
}
match provider.search(&query, &opts).await {
Ok(results) => (name, results),
Err(e) => {
tracing::error!("Provider {} failed: {}", name, e);
(name, Vec::new())
}
}
}));
}
let mut results_by_provider = HashMap::new();
for handle in handles {
if let Ok((name, results)) = handle.await {
results_by_provider.insert(name, results);
}
}
Ok(merge_results(&results_by_provider, &merge_opts))
}
pub async fn search_single(
&self,
query: &str,
provider_name: &str,
options: SearchOptions,
) -> Result<Vec<SearchResult>, SearchError> {
let provider = self
.providers
.get(provider_name)
.ok_or_else(|| SearchError::UnknownProvider(provider_name.to_string()))?;
let provider = provider.read().await;
if !provider.is_available() {
return Err(SearchError::ProviderDisabled(provider_name.to_string()));
}
provider.search(query, &options).await
}
pub fn get_available_providers(&self) -> Vec<String> {
self.providers.keys().cloned().collect()
}
pub fn get_registry(&self) -> &[RegistryEntry] {
&self.registry
}
pub async fn get_provider_status(&self) -> HashMap<String, ProviderStatus> {
let mut status = HashMap::new();
for (name, provider) in &self.providers {
let p = provider.read().await;
let meta = self.registry.iter().find(|e| &e.id == name);
status.insert(
name.clone(),
ProviderStatus {
enabled: p.is_available(),
weight: p.weight(),
category: meta.map(|m| m.category.clone()),
label: meta.map(|m| m.label.clone()),
cors_readable: meta.map(|m| m.cors_readable),
access: meta.map(|m| m.access.clone()),
},
);
}
status
}
pub async fn set_provider_weight(&self, name: &str, weight: f64) -> Result<(), SearchError> {
let provider = self
.providers
.get(name)
.ok_or_else(|| SearchError::UnknownProvider(name.to_string()))?;
provider.write().await.set_weight(weight);
Ok(())
}
pub async fn set_provider_enabled(&self, name: &str, enabled: bool) -> Result<(), SearchError> {
let provider = self
.providers
.get(name)
.ok_or_else(|| SearchError::UnknownProvider(name.to_string()))?;
provider.write().await.set_enabled(enabled);
Ok(())
}
}
impl Default for WebSearchEngine {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, serde::Serialize)]
#[serde(rename_all = "camelCase")]
pub struct ProviderStatus {
pub enabled: bool,
pub weight: f64,
#[serde(skip_serializing_if = "Option::is_none")]
pub category: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub label: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cors_readable: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub access: Option<String>,
}