use axum::Json;
use axum::extract::State;
use axum::http::HeaderMap;
use serde::{Deserialize, Serialize};
mod brave;
use crate::AppState;
use crate::check_auth;
use crate::config::{Secret, WebSearchConfig};
use crate::error::GatewayError;
use crate::web_search_process::post_process_results;
use brave::{BraveSearchParams, brave_overfetch_count, brave_search};
#[derive(Debug, Clone)]
pub(crate) struct WebSearchSettings {
pub default_count: u8,
pub max_count: u8,
pub max_per_host: u8,
pub default_freshness: String,
pub default_safesearch: String,
pub strip_tracking: bool,
}
impl WebSearchSettings {
#[must_use]
pub(crate) fn from_config(cfg: &WebSearchConfig) -> WebSearchSettings {
WebSearchSettings {
default_count: cfg.default_count,
max_count: cfg.max_count,
max_per_host: cfg.max_per_host,
default_freshness: cfg.default_freshness.clone(),
default_safesearch: cfg.default_safesearch.clone(),
strip_tracking: cfg.strip_tracking,
}
}
}
#[derive(Debug)]
pub(crate) struct WebSearchState {
api_key: Secret,
base_url: String,
pub settings: WebSearchSettings,
http: reqwest::Client,
}
impl WebSearchState {
#[must_use]
pub(crate) fn new(cfg: &WebSearchConfig) -> WebSearchState {
let crate::config::SearchProvider::Brave = cfg.provider;
WebSearchState {
api_key: cfg.api_key.clone(),
base_url: cfg.base_url.trim_end_matches('/').to_string(),
settings: WebSearchSettings::from_config(cfg),
http: crate::http_util::bounded_client(),
}
}
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub(crate) struct WebSearchRequest {
pub query: String,
#[serde(default)]
pub count: Option<u8>,
#[serde(default)]
pub freshness: Option<String>,
#[serde(default)]
pub country: Option<String>,
#[serde(default)]
pub search_lang: Option<String>,
#[serde(default)]
pub safesearch: Option<String>,
#[serde(default)]
pub include_domains: Vec<String>,
#[serde(default)]
pub exclude_domains: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub(crate) struct WebSearchResponse {
pub query: String,
pub results: Vec<SearchResult>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub(crate) struct SearchResult {
pub title: String,
pub url: String,
pub description: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub age: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub site_name: Option<String>,
#[serde(skip_serializing_if = "Vec::is_empty")]
pub extra_snippets: Vec<String>,
}
const MAX_QUERY_CHARS: usize = 512;
fn trim_web_search_query(query: &str) -> Result<String, GatewayError> {
let trimmed = query.trim();
if trimmed.is_empty() {
return Err(GatewayError::MalformedRequest(
"web_search: empty query".to_string(),
));
}
Ok(trimmed.chars().take(MAX_QUERY_CHARS).collect())
}
fn validate_domain_filters(field: &str, domains: &[String]) -> Result<Vec<String>, GatewayError> {
domains
.iter()
.map(|raw| validate_domain_filter(field, raw))
.collect()
}
fn validate_domain_filter(field: &str, raw: &str) -> Result<String, GatewayError> {
let domain = raw.trim();
let malformed =
|| GatewayError::MalformedRequest(format!("web_search: invalid {field} domain {raw:?}"));
if domain.is_empty()
|| domain.len() > 253
|| domain.contains("://")
|| domain.contains('/')
|| domain.contains(':')
|| domain.chars().any(char::is_whitespace)
{
return Err(malformed());
}
let lower = domain.to_ascii_lowercase();
if !is_valid_domain_syntax(&lower) {
return Err(malformed());
}
Ok(lower)
}
fn is_valid_domain_syntax(domain: &str) -> bool {
let mut labels = 0_usize;
for label in domain.split('.') {
if label.is_empty()
|| label.len() > 63
|| label.starts_with('-')
|| label.ends_with('-')
|| !label
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'-')
{
return false;
}
labels += 1;
}
labels >= 1
}
#[must_use]
fn clamp_count(requested: u8, max_count: u8) -> u8 {
let max_count = max_count.max(1);
requested.clamp(1, max_count)
}
fn validate_request_knobs(request: &WebSearchRequest) -> Result<(), GatewayError> {
if let Some(freshness) = non_empty_opt(request.freshness.as_deref())
&& !is_valid_freshness(freshness)
{
return Err(GatewayError::MalformedRequest(format!(
"web_search: invalid freshness {freshness:?}"
)));
}
if let Some(safesearch) = non_empty_opt(request.safesearch.as_deref())
&& !matches!(safesearch, "off" | "moderate" | "strict")
{
return Err(GatewayError::MalformedRequest(format!(
"web_search: invalid safesearch {safesearch:?}"
)));
}
if let Some(country) = non_empty_opt(request.country.as_deref())
&& !is_alpha_code(country, 2, 2)
{
return Err(GatewayError::MalformedRequest(format!(
"web_search: invalid country {country:?}"
)));
}
if let Some(lang) = non_empty_opt(request.search_lang.as_deref())
&& !is_alpha_code(lang, 2, 3)
{
return Err(GatewayError::MalformedRequest(format!(
"web_search: invalid search_lang {lang:?}"
)));
}
Ok(())
}
fn is_valid_freshness(value: &str) -> bool {
if matches!(value, "pd" | "pw" | "pm" | "py") {
return true;
}
value
.split_once("to")
.is_some_and(|(from, to)| is_iso_date(from) && is_iso_date(to))
}
fn is_iso_date(value: &str) -> bool {
let bytes = value.as_bytes();
bytes.len() == 10
&& bytes[4] == b'-'
&& bytes[7] == b'-'
&& bytes
.iter()
.enumerate()
.all(|(index, byte)| index == 4 || index == 7 || byte.is_ascii_digit())
}
fn is_alpha_code(value: &str, min: usize, max: usize) -> bool {
let len = value.chars().count();
len >= min && len <= max && value.chars().all(|c| c.is_ascii_alphabetic())
}
fn non_empty_opt(value: Option<&str>) -> Option<&str> {
value.map(str::trim).filter(|s| !s.is_empty())
}
fn resolve_freshness<'a>(request: Option<&'a str>, default_freshness: &'a str) -> Option<&'a str> {
non_empty_opt(request).or_else(|| non_empty_opt(Some(default_freshness)))
}
fn resolve_safesearch<'a>(
request: Option<&'a str>,
default_safesearch: &'a str,
) -> Option<&'a str> {
non_empty_opt(request).or_else(|| non_empty_opt(Some(default_safesearch)))
}
pub(crate) async fn web_search(
State(state): State<AppState>,
headers: HeaderMap,
Json(request): Json<WebSearchRequest>,
) -> Result<Json<WebSearchResponse>, GatewayError> {
check_auth(&state, &headers).await?;
let web_search = state
.web_search()
.await
.ok_or(GatewayError::ToolNotConfigured("web_search"))?;
let query = trim_web_search_query(&request.query)?;
validate_request_knobs(&request)?;
let include_domains = validate_domain_filters("include", &request.include_domains)?;
let exclude_domains = validate_domain_filters("exclude", &request.exclude_domains)?;
let count = clamp_count(
request.count.unwrap_or(web_search.settings.default_count),
web_search.settings.max_count,
);
let brave_count = brave_overfetch_count(count, web_search.settings.max_count);
let params = BraveSearchParams {
query: &query,
count: brave_count,
freshness: resolve_freshness(
request.freshness.as_deref(),
&web_search.settings.default_freshness,
),
country: non_empty_opt(request.country.as_deref()),
search_lang: non_empty_opt(request.search_lang.as_deref()),
safesearch: resolve_safesearch(
request.safesearch.as_deref(),
&web_search.settings.default_safesearch,
),
};
let mapped = brave_search(
&web_search.http,
&web_search.base_url,
web_search.api_key.expose(),
¶ms,
)
.await?;
let results = post_process_results(
mapped,
web_search.settings.strip_tracking,
&include_domains,
&exclude_domains,
web_search.settings.max_per_host,
count,
);
Ok(Json(WebSearchResponse { query, results }))
}
#[cfg(test)]
mod tests {
use super::brave::{brave_search_query, prefix_web_search_upstream};
use super::*;
use crate::error::GatewayError;
#[test]
fn empty_query_is_malformed_request() {
for query in ["", " ", "\t\n"] {
let err = trim_web_search_query(query).expect_err("empty query");
match err {
GatewayError::MalformedRequest(message) => {
assert_eq!(message, "web_search: empty query");
}
other => panic!("expected MalformedRequest, got {other:?}"),
}
}
}
fn knob_request(
freshness: &str,
safesearch: &str,
country: &str,
lang: &str,
) -> WebSearchRequest {
WebSearchRequest {
query: "q".to_string(),
count: None,
freshness: Some(freshness.to_string()),
country: Some(country.to_string()),
search_lang: Some(lang.to_string()),
safesearch: Some(safesearch.to_string()),
include_domains: Vec::new(),
exclude_domains: Vec::new(),
}
}
#[test]
fn validate_request_knobs_accepts_valid_and_empty() {
assert!(validate_request_knobs(&knob_request("pd", "moderate", "us", "en")).is_ok());
assert!(
validate_request_knobs(&knob_request("2024-01-01to2024-12-31", "off", "GB", "eng"))
.is_ok()
);
assert!(validate_request_knobs(&knob_request("", "", "", "")).is_ok());
}
#[test]
fn validate_request_knobs_rejects_malformed() {
for req in [
knob_request("daily", "", "", ""),
knob_request("", "medium", "", ""),
knob_request("", "", "usa", ""),
knob_request("", "", "", "english"),
knob_request("", "", "1", ""),
] {
assert!(matches!(
validate_request_knobs(&req),
Err(GatewayError::MalformedRequest(_))
));
}
}
#[test]
fn validate_domain_filters_accepts_valid_and_rejects_malformed() {
assert_eq!(
validate_domain_filters("include", &["Example.COM".into(), "sub.a-b.co".into()])
.expect("valid domains"),
vec!["example.com".to_string(), "sub.a-b.co".to_string()]
);
assert!(
validate_domain_filters("exclude", &[])
.expect("empty list")
.is_empty()
);
for bad in [
"", " ", "https://example.com", "example.com/path", "exa mple.com", "exa$mple.com", "example.com:8080", "-bad.com", "bad-.com", "a..b.com", ] {
let err = validate_domain_filters("include", &[bad.to_string()])
.expect_err(&format!("{bad:?} must be rejected"));
assert!(matches!(err, GatewayError::MalformedRequest(_)), "{err:?}");
}
}
#[test]
fn non_empty_query_is_trimmed() {
assert_eq!(
trim_web_search_query(" C++ Alliance ").expect("ok"),
"C++ Alliance"
);
}
#[test]
fn brave_overfetch_uses_triple_capped_by_max() {
assert_eq!(brave_overfetch_count(5, 20), 15);
assert_eq!(brave_overfetch_count(10, 20), 20);
assert_eq!(brave_overfetch_count(1, 20), 3);
assert_eq!(brave_overfetch_count(20, 20), 20);
}
#[test]
fn clamp_count_bounds_to_one_through_max() {
assert_eq!(clamp_count(0, 20), 1);
assert_eq!(clamp_count(5, 20), 5);
assert_eq!(clamp_count(50, 20), 20);
}
#[test]
fn resolve_knobs_prefer_request_then_defaults() {
assert_eq!(resolve_freshness(Some("pd"), "pw"), Some("pd"));
assert_eq!(resolve_freshness(Some(""), "pw"), Some("pw"));
assert_eq!(resolve_freshness(None, ""), None);
assert_eq!(resolve_safesearch(None, "moderate"), Some("moderate"));
assert_eq!(non_empty_opt(Some("us")), Some("us"));
assert_eq!(non_empty_opt(Some(" ")), None);
}
#[test]
fn prefix_web_search_upstream_prefixes_status_body() {
let err = prefix_web_search_upstream(GatewayError::UpstreamStatus {
status: 429,
body: "rate limited".to_string(),
});
match err {
GatewayError::UpstreamStatus { body, .. } => {
assert_eq!(body, "web_search: rate limited");
}
other => panic!("expected UpstreamStatus, got {other:?}"),
}
}
#[test]
fn brave_search_query_always_sets_extra_snippets_and_optional_knobs() {
let base = BraveSearchParams {
query: "C++ Alliance",
count: 15,
freshness: None,
country: None,
search_lang: None,
safesearch: None,
};
let pairs = brave_search_query(&base);
assert_eq!(
pairs,
vec![
("q", "C++ Alliance".to_string()),
("count", "15".to_string()),
("extra_snippets", "true".to_string()),
]
);
let full = BraveSearchParams {
query: "boost",
count: 9,
freshness: Some("pd"),
country: Some("us"),
search_lang: Some("en"),
safesearch: Some("moderate"),
};
let pairs = brave_search_query(&full);
assert_eq!(
pairs,
vec![
("q", "boost".to_string()),
("count", "9".to_string()),
("extra_snippets", "true".to_string()),
("freshness", "pd".to_string()),
("country", "us".to_string()),
("search_lang", "en".to_string()),
("safesearch", "moderate".to_string()),
]
);
}
}
#[cfg(test)]
mod live_tests {
use super::brave::{BraveSearchParams, brave_search};
#[tokio::test]
#[ignore = "hits the real Brave API; requires BRAVE_API_KEY, run with --ignored"]
async fn live_brave_search() {
let api_key =
std::env::var("BRAVE_API_KEY").expect("set BRAVE_API_KEY to run this live test");
let http = reqwest::Client::new();
let params = BraveSearchParams {
query: "rust programming language",
count: 5,
freshness: None,
country: None,
search_lang: None,
safesearch: None,
};
let results = brave_search(
&http,
"https://api.search.brave.com/res/v1",
&api_key,
¶ms,
)
.await
.expect("brave search should succeed");
assert!(
!results.is_empty(),
"expected at least one result from Brave"
);
assert!(
results[0].url.starts_with("http"),
"expected a real URL, got: {}",
results[0].url
);
for result in &results {
println!("- {} :: {}", result.title, result.url);
}
}
}