diff --git a/Cargo.lock b/Cargo.lock index 5308f283..44e43e78 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -97,6 +97,7 @@ dependencies = [ "serde_yml", "sqlx", "sse-stream", + "strum", "thiserror", "tokio", "tokio-util", @@ -3441,6 +3442,27 @@ version = "0.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" +[[package]] +name = "strum" +version = "0.27.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "af23d6f6c1a224baef9d3f61e287d2761385a5b88fdab4eb4c6f11aeb54c4bcf" +dependencies = [ + "strum_macros", +] + +[[package]] +name = "strum_macros" +version = "0.27.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7695ce3845ea4b33927c055a39dc438a45b059f7c1b3d91d38d10355fb8cbca7" +dependencies = [ + "heck", + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "subtle" version = "2.6.1" diff --git a/Cargo.toml b/Cargo.toml index ae9f698b..8330a431 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -36,6 +36,9 @@ serde = { version = "1", features = ["derive"] } # Structured Outputs emit object keys in JSON Schema declaration order. serde_json = { version = "1", features = ["preserve_order", "raw_value"] } sse-stream = "0.2.3" +# Derives only the variant-name list for configuration error messages, keeping +# it in lockstep with serde's accepted wire names (WebSearchProviderKind::parse_name). +strum = { version = "0.27", features = ["derive"] } thiserror = "2" tokio = { version = "1", features = ["full"] } tokio-util = "0.7" diff --git a/README.md b/README.md index f47fa300..050629ec 100644 --- a/README.md +++ b/README.md @@ -219,6 +219,9 @@ llm_api_base = "http://127.0.0.1:5050" # database_url = "postgresql://agentic-api@localhost/agentic_api" [web_search] +# Backend that serves the gateway-owned `web_search` tool: "you" (default) or +# "brave". Overridden by AGENTIC_WEB_SEARCH_PROVIDER. +provider = "you" base_url = "https://api.ydc-index.io" api_key_env = "YOU_API_KEY" @@ -253,8 +256,19 @@ parsed. Order of precedence is `--max-request-body-size-bytes`, then `AGENTIC_MA file setting. `api_key_env` names the process environment variable containing the web-search credential; it does not contain the -credential itself. `YOU_API_BASE_URL`, `AGENTIC_MCP_ALLOWED_HOSTS`, `AGENTIC_MAX_REQUEST_BODY_SIZE_BYTES`, and -`AGENTIC_MAX_CONCURRENT_GATEWAY_CALLS` can override their typed file settings. The concurrency value is a sliding-window +credential itself. When `api_key_env` is unset, the selected provider's default name is used: `YOU_API_KEY` for +`"you"` and `BRAVE_API_KEY` for `"brave"`, so switching providers never requires editing the configuration file. +`AGENTIC_WEB_SEARCH_PROVIDER`, `AGENTIC_WEB_SEARCH_BASE_URL`, +`YOU_API_BASE_URL`, `AGENTIC_MCP_ALLOWED_HOSTS`, `AGENTIC_MAX_REQUEST_BODY_SIZE_BYTES`, and +`AGENTIC_MAX_CONCURRENT_GATEWAY_CALLS` can override their typed file settings. +`AGENTIC_WEB_SEARCH_BASE_URL` applies to the alternative `brave` provider; You.com +keeps its historical `YOU_API_BASE_URL` override. The `brave` provider (Brave +Search API) is the default alternative to You.com: it clamps `count` to its 20-result +per-section cap, has no server-side domain filtering (domain allow/block lists are +post-filtered client-side), and runs at most one request in flight to respect the +free-tier ~1 QPS rate limit. One consequence of client-side filtering: You.com rejects a request that combines +`include_domains` with `exclude_domains`, while `brave` accepts the combination and applies the blocklist on top of +the allowlist. The concurrency value is a sliding-window upper bound; handlers may further serialize calls to the same tool name. The MCP allowlist is used only for request-declared remote MCP URLs; configured `[mcp_servers]` entries are trusted operator configuration. diff --git a/crates/agentic-server-core/Cargo.toml b/crates/agentic-server-core/Cargo.toml index 543eb24f..0446110c 100644 --- a/crates/agentic-server-core/Cargo.toml +++ b/crates/agentic-server-core/Cargo.toml @@ -40,6 +40,7 @@ rmcp = { workspace = true, features = [ serde.workspace = true serde_json.workspace = true sse-stream.workspace = true +strum.workspace = true thiserror.workspace = true tokio = { workspace = true, features = ["time"] } tokio-util = { workspace = true, features = ["rt"] } diff --git a/crates/agentic-server-core/src/config.rs b/crates/agentic-server-core/src/config.rs index 10ffd711..279aa069 100644 --- a/crates/agentic-server-core/src/config.rs +++ b/crates/agentic-server-core/src/config.rs @@ -90,23 +90,61 @@ impl Default for SqliteConfig { /// Backend that serves the gateway-owned `web_search` tool. /// /// Additional providers are added here (#291). The enum is non-exhaustive so -/// downstream crates keep a fallback arm when a new variant lands. Selecting a -/// provider through [`WebSearchProviderConfig`] is deferred until a second -/// provider exists. +/// downstream crates keep a fallback arm when a new variant lands. #[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] #[non_exhaustive] +#[derive(strum::EnumIter, strum::VariantNames)] +#[strum(serialize_all = "snake_case")] pub enum WebSearchProviderKind { #[default] You, + Brave, } impl WebSearchProviderKind { + /// Every provider kind's `snake_case` name, in declaration order. Powers + /// configuration error messages without duplicating the variant list. + pub const VARIANTS: &'static [&'static str] = ::VARIANTS; + + /// Iterates over every provider kind. + pub fn variants() -> impl Iterator { + ::iter() + } + + /// Parses a provider name using the same `snake_case` wire names serde + /// accepts in configuration files, so environment parsing and file parsing + /// share one set of accepted names. Environment values are trimmed and + /// case-insensitive, matching operator expectations; the error carries the + /// caller's context. + /// + /// This deliberately round-trips through serde instead of strum's + /// `EnumString`: `FromStr` would derive a parallel name set that only a + /// test (not the type system) keeps in sync with serde's, whereas routing + /// through serde makes file and environment parsing consistent by + /// construction. + /// + /// # Errors + /// + /// Returns [`crate::error::Error::Config`] when `value` (trimmed, + /// case-insensitively) is not a `snake_case` name of a known provider + /// variant. + pub fn parse_name(context: &str, value: &str) -> Result { + let normalized = value.trim().to_ascii_lowercase(); + serde_json::from_value::(serde_json::Value::String(normalized)).map_err(|_| { + crate::error::Error::Config(format!( + "invalid {context} value '{value}': expected one of '{}'", + Self::VARIANTS.join("', '") + )) + }) + } + /// Environment variable that conventionally carries this provider's API key. #[must_use] pub const fn default_api_key_env(self) -> &'static str { match self { Self::You => "YOU_API_KEY", + Self::Brave => "BRAVE_API_KEY", } } @@ -115,6 +153,7 @@ impl WebSearchProviderKind { pub const fn display_name(self) -> &'static str { match self { Self::You => "You.com", + Self::Brave => "Brave", } } } @@ -125,18 +164,27 @@ impl std::fmt::Display for WebSearchProviderKind { } } -/// Credentials for the gateway-owned `web_search` provider (You.com). +/// Credentials and selection for the gateway-owned `web_search` provider. +/// +/// `kind` chooses the backend; the credential and endpoint are resolved per +/// provider at deployment time. #[derive(Clone, Default)] pub struct WebSearchProviderConfig { + pub kind: WebSearchProviderKind, pub api_key: Option, pub base_url: Option, } impl WebSearchProviderConfig { - /// Builds the config from the credential and endpoint the deployment resolved. + /// Builds the config from the selection and credential the deployment + /// resolved. #[must_use] - pub const fn new(api_key: Option, base_url: Option) -> Self { - Self { api_key, base_url } + pub fn new(kind: WebSearchProviderKind, api_key: Option, base_url: Option) -> Self { + Self { + kind, + api_key, + base_url, + } } } @@ -144,6 +192,7 @@ impl std::fmt::Debug for WebSearchProviderConfig { /// Redacts `api_key` so debug-printing any enclosing config never logs the secret. fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("WebSearchProviderConfig") + .field("kind", &self.kind) .field("api_key", &self.api_key.as_ref().map(|_| "")) .field("base_url", &self.base_url) .finish() @@ -300,6 +349,7 @@ mod tests { #[test] fn web_search_provider_config_debug_redacts_api_key() { let config = WebSearchProviderConfig::new( + WebSearchProviderKind::You, Some("super-secret-key".to_owned()), Some("https://api.example".to_owned()), ); @@ -315,7 +365,7 @@ mod tests { assert!(!format!("{tools:?}").contains("super-secret-key")); assert_eq!( format!("{:?}", WebSearchProviderConfig::default()), - "WebSearchProviderConfig { api_key: None, base_url: None }" + "WebSearchProviderConfig { kind: You, api_key: None, base_url: None }" ); } @@ -323,9 +373,59 @@ mod tests { fn web_search_provider_kind_labels() { assert_eq!(WebSearchProviderKind::You.to_string(), "You.com"); assert_eq!(WebSearchProviderKind::You.default_api_key_env(), "YOU_API_KEY"); + assert_eq!(WebSearchProviderKind::Brave.to_string(), "Brave"); + assert_eq!(WebSearchProviderKind::Brave.default_api_key_env(), "BRAVE_API_KEY"); assert_eq!(WebSearchProviderKind::default(), WebSearchProviderKind::You); } + #[test] + fn web_search_provider_kind_variants_stay_in_sync() { + // strum's variant list and iterator must match serde's accepted wire + // names so configuration errors and parsing never drift apart. + assert_eq!(WebSearchProviderKind::VARIANTS, ["you", "brave"]); + let names: Vec = WebSearchProviderKind::variants().map(|kind| kind.to_string()).collect(); + assert_eq!(names, ["You.com", "Brave"]); + for name in WebSearchProviderKind::VARIANTS { + let parsed: WebSearchProviderKind = + serde_json::from_value(serde_json::Value::String((*name).to_owned())).expect("serde parses variant"); + // Round-tripping the strum-derived name through serde proves the + // two derive the same set of accepted wire names. + let rendered = serde_json::to_value(parsed).expect("serde serializes variant"); + assert_eq!(rendered, serde_json::Value::String((*name).to_owned())); + } + } + + #[test] + fn web_search_provider_kind_parse_name_is_trimmed_and_case_insensitive() { + // Environment values are trimmed and case-insensitive, so common + // operator spellings such as `Brave` or ` brave ` are accepted. + assert_eq!( + WebSearchProviderKind::parse_name("env", "brave").expect("parses"), + WebSearchProviderKind::Brave + ); + assert_eq!( + WebSearchProviderKind::parse_name("env", "Brave").expect("parses"), + WebSearchProviderKind::Brave + ); + assert_eq!( + WebSearchProviderKind::parse_name("env", " BRAVE ").expect("parses"), + WebSearchProviderKind::Brave + ); + assert_eq!( + WebSearchProviderKind::parse_name("env", "you").expect("parses"), + WebSearchProviderKind::You + ); + + // Unknown names are rejected, and the error quotes the original value. + let error = WebSearchProviderKind::parse_name("TEST_PROVIDER", "nope").expect_err("rejects unknown"); + assert!( + error + .to_string() + .contains("invalid TEST_PROVIDER value 'nope': expected one of 'you', 'brave'"), + "unexpected error message: {error}" + ); + } + #[test] fn strip_trailing_v1() { assert_eq!(normalize_base_url("http://host:8000/v1"), "http://host:8000"); diff --git a/crates/agentic-server-core/src/tool/executors.rs b/crates/agentic-server-core/src/tool/executors.rs index 02e63352..2cf9ff7f 100644 --- a/crates/agentic-server-core/src/tool/executors.rs +++ b/crates/agentic-server-core/src/tool/executors.rs @@ -84,8 +84,9 @@ impl GatewayExecutors { } else { config.mcp_allowed_hosts.clone() }, - web_search: Some(Arc::new(WebSearchHandler::from_values( + web_search: Some(Arc::new(WebSearchHandler::from_provider( client, + config.web_search.kind, config.web_search.api_key.clone(), config.web_search.base_url.clone(), config.max_concurrent_gateway_calls, diff --git a/crates/agentic-server-core/src/tool/web_search/args.rs b/crates/agentic-server-core/src/tool/web_search/args.rs index e1f3214f..3052efcc 100644 --- a/crates/agentic-server-core/src/tool/web_search/args.rs +++ b/crates/agentic-server-core/src/tool/web_search/args.rs @@ -6,6 +6,7 @@ //! filtering. use std::fmt; +use std::num::NonZeroUsize; use std::str::FromStr; use chrono::NaiveDate; @@ -205,22 +206,19 @@ pub(crate) fn clean_vec(values: Option<&[String]>) -> Option> { } /// Provider-neutral domain post-filter for providers without server-side -/// `include_domains` / `exclude_domains` support. +/// `include_domains` / `exclude_domains` support (e.g. Brave). /// /// A host matches a domain when it equals the domain or ends with `.{domain}` /// (label boundary), compared case-insensitively after IDNA normalization. A /// URL without a parseable host cannot be checked, so it is rejected whenever /// any allowlist or blocklist is active (fail closed). You.com filters -/// server-side, so this is not applied on that path; the first provider that -/// needs it wires it in (#291 Phase 2). +/// server-side, so this is not applied on that path. #[derive(Debug, Clone, Default, PartialEq, Eq)] -#[allow(dead_code)] // wired by the first provider without server-side filtering (#291 Phase 2) pub(crate) struct DomainFilter { include: Vec, exclude: Vec, } -#[allow(dead_code)] // wired by the first provider without server-side filtering (#291 Phase 2) impl DomainFilter { pub(crate) fn new(include: Option<&[String]>, exclude: Option<&[String]>) -> Self { Self { @@ -278,10 +276,32 @@ fn host_matches_domain(host: &str, domain: &str) -> bool { host == domain || host.strip_suffix(domain).is_some_and(|prefix| prefix.ends_with('.')) } +/// Caps the requested query concurrency at a provider's own ceiling. +/// +/// `None` leaves the requested value untouched; otherwise the result is +/// `min(requested, ceiling)`. +pub(crate) fn cap_provider_concurrency(ceiling: Option, requested: NonZeroUsize) -> NonZeroUsize { + ceiling.map_or(requested, |cap| requested.min(cap)) +} + #[cfg(test)] mod tests { use super::*; + #[test] + fn provider_concurrency_cap_bounds_the_request() { + let requested = NonZeroUsize::new(5).unwrap(); + assert_eq!(cap_provider_concurrency(None, requested), requested); + assert_eq!( + cap_provider_concurrency(Some(NonZeroUsize::new(2).unwrap()), requested), + NonZeroUsize::new(2).unwrap() + ); + assert_eq!( + cap_provider_concurrency(Some(NonZeroUsize::new(8).unwrap()), requested), + requested + ); + } + fn result(url: &str) -> WebSearchResult { WebSearchResult { url: url.to_owned(), diff --git a/crates/agentic-server-core/src/tool/web_search/brave.rs b/crates/agentic-server-core/src/tool/web_search/brave.rs new file mode 100644 index 00000000..d8009d05 --- /dev/null +++ b/crates/agentic-server-core/src/tool/web_search/brave.rs @@ -0,0 +1,498 @@ +//! Brave Search API provider for `web_search`. +//! +//! Owns request shaping against Brave's `GET /res/v1/web/search` and the +//! mapping of its JSON envelope onto the provider-neutral +//! [`WebSearchProviderResponse`]. +//! +//! Unlike You.com, Brave has no server-side domain filtering: `include_domains` +//! and `exclude_domains` are post-filtered client-side by host suffix on a +//! label boundary. `count` is capped at 20 (clamped, never an error) and the +//! free-tier rate limit is ~1 QPS, so the provider's own concurrency ceiling is +//! 1. +//! +//! Transport rule: the gateway's `reqwest` client is built without gzip +//! support (see `crates/agentic-server-core/Cargo.toml`), so no +//! `Accept-Encoding: gzip` header is sent or expected. + +use std::fmt; +use std::future::Future; +use std::num::NonZeroUsize; +use std::pin::Pin; +use std::sync::Arc; + +use serde::Deserialize; + +use super::args::{DomainFilter, Freshness, WebSearchArguments, clean_string, clean_vec, validate_count}; +use super::{ + WebSearchProvider, WebSearchProviderMetadata, WebSearchProviderResponse, WebSearchResult, null_as_default, + read_response_limited, +}; +use crate::config::WebSearchProviderKind; +use crate::tool::handler::ToolError; +use crate::types::tools::{WebSearchContextSize, WebSearchToolParam}; + +pub(crate) const BRAVE_API_KEY: &str = WebSearchProviderKind::Brave.default_api_key_env(); +/// Default Brave Search API base URL used when no base URL is configured. The +/// deployment-facing override is `AGENTIC_WEB_SEARCH_BASE_URL`, resolved by the +/// server binary before a provider is built; unlike You.com, Brave has a +/// usable default endpoint. +pub(crate) const BRAVE_DEFAULT_BASE_URL: &str = "https://api.search.brave.com"; +/// Free-tier developer cap on results per section; requested counts above it +/// are clamped rather than rejected, since the model cannot predict provider +/// limits. +const BRAVE_MAX_COUNT: u8 = 20; +/// Free-tier rate limit of ~1 QPS: the provider never runs more than one +/// request in flight. +const BRAVE_MAX_CONCURRENT_REQUESTS: NonZeroUsize = NonZeroUsize::new(1).expect("brave ceiling is nonzero"); + +/// Provider credential whose `Debug` output never contains the secret. +#[derive(Clone)] +pub(crate) struct ApiKey(pub String); + +impl fmt::Debug for ApiKey { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("ApiKey()") + } +} + +#[derive(Debug, Clone)] +pub(crate) struct BraveSearchProvider { + client: Arc, + api_key: Option, + base_url: Option, +} + +impl BraveSearchProvider { + /// Builds a provider from optional environment-style values: a blank key + /// counts as unset and fails at execution time. A blank base URL falls + /// back to [`BRAVE_DEFAULT_BASE_URL`]. + pub(crate) fn from_values(client: Arc, api_key: Option, base_url: Option) -> Self { + let api_key = api_key + .map(|value| value.trim().to_owned()) + .filter(|value| !value.is_empty()) + .map(ApiKey); + let base_url = base_url + .and_then(|value| clean_base_url(&value)) + .or_else(|| Some(BRAVE_DEFAULT_BASE_URL.to_owned())); + Self { + client, + api_key, + base_url, + } + } +} + +impl WebSearchProvider for BraveSearchProvider { + fn search<'a>( + &'a self, + query: &'a str, + args: &'a WebSearchArguments, + config: &'a WebSearchToolParam, + ) -> Pin> + Send + 'a>> { + Box::pin(async move { + let api_key = self + .api_key + .as_ref() + .ok_or_else(|| ToolError::Config(format!("{BRAVE_API_KEY} must be set to use the web_search tool")))?; + // `from_values` falls back to `BRAVE_DEFAULT_BASE_URL`, so the base + // URL is always set for a provider built through the constructor. + let base_url = self.base_url.as_deref().unwrap_or(BRAVE_DEFAULT_BASE_URL); + let request = BraveSearchRequest::from_args_and_config(query, args, config)?; + let resp = self + .client + .get(format!("{base_url}/res/v1/web/search")) + .query(&request.query_params()) + .header("X-Subscription-Token", &api_key.0) + .send() + .await + .map_err(|e| ToolError::Execution(format!("Brave search request failed: {e}")))?; + + if !resp.status().is_success() { + let status = resp.status(); + // Surface the upstream rate-limit hint so operators can back off; + // no automatic retry (Phase 2 scope). + let retry_after = resp + .headers() + .get("retry-after") + .and_then(|value| value.to_str().ok()) + .map(str::to_owned); + let body = read_response_limited(resp, WebSearchProviderKind::Brave) + .await + .unwrap_or_default(); + let retry_note = retry_after + .map(|value| format!(" (retry after {value})")) + .unwrap_or_default(); + return Err(ToolError::Execution(format!( + "Brave search returned {status}{retry_note}: {body}" + ))); + } + + let response_text = read_response_limited(resp, WebSearchProviderKind::Brave).await?; + let response: BraveSearchResponse = serde_json::from_str(&response_text) + .map_err(|e| ToolError::Execution(format!("Brave search returned invalid JSON: {e}")))?; + Ok(response.into_provider_response(&request.query, &request.domain_filter)) + }) + } + + fn max_concurrent_requests(&self) -> Option { + Some(BRAVE_MAX_CONCURRENT_REQUESTS) + } +} + +/// Query parameters for Brave's `GET /res/v1/web/search`, derived from the +/// model's arguments and the request-level tool configuration. +#[derive(Debug, PartialEq)] +struct BraveSearchRequest { + query: String, + count: Option, + freshness: Option, + country: Option, + language: Option, + domain_filter: DomainFilter, +} + +impl BraveSearchRequest { + fn query_params(&self) -> Vec<(String, String)> { + let mut params = vec![ + ("q".to_owned(), self.query.clone()), + ("result_filter".to_owned(), "web,news".to_owned()), + ]; + if let Some(count) = self.count { + params.push(("count".to_owned(), count.to_string())); + } + if let Some(freshness) = &self.freshness { + let rendered = brave_freshness(freshness); + if !rendered.is_empty() { + params.push(("freshness".to_owned(), rendered)); + } + } + if let Some(country) = &self.country { + params.push(("country".to_owned(), country.clone())); + } + if let Some(language) = &self.language { + params.push(("search_lang".to_owned(), language.clone())); + } + params + } + + fn from_args_and_config( + query: &str, + args: &WebSearchArguments, + config: &WebSearchToolParam, + ) -> Result { + let count = args + .count + .or_else(|| { + config + .search_context_size + .map(WebSearchContextSize::default_count) + .map(u16::from) + }) + .map(clamp_count) + .transpose()?; + let config_domains = config + .filters + .as_ref() + .and_then(|filters| clean_vec(filters.allowed_domains.as_deref())); + let config_blocked_domains = config + .filters + .as_ref() + .and_then(|filters| clean_vec(filters.blocked_domains.as_deref())); + let include_domains = config_domains.or_else(|| args.include_domains.clone()); + let exclude_domains = config_blocked_domains.or_else(|| args.exclude_domains.clone()); + // Brave cannot filter domains server-side; remember them for a + // client-side post-filter instead of dropping or erroring. + let domain_filter = DomainFilter::new(include_domains.as_deref(), exclude_domains.as_deref()); + if args.boost_domains.is_some() { + tracing::debug!("web_search boost_domains ignored: Brave does not support boosting"); + } + if args.livecrawl.is_some() || args.livecrawl_formats.is_some() || args.crawl_timeout.is_some() { + tracing::debug!("web_search livecrawl arguments ignored: You.com-specific, unsupported by Brave"); + } + let country = config + .user_location + .as_ref() + .and_then(|location| clean_string(location.country.as_deref())) + .or_else(|| args.country.clone()) + .map(|value| value.to_ascii_uppercase()); + + Ok(Self { + query: query.trim().to_owned(), + count, + freshness: args.freshness, + country, + language: args.language.clone(), + domain_filter, + }) + } +} + +/// Renders a typed freshness filter in Brave's syntax (`pd`/`pw`/`pm`/`py`). +/// +/// A custom date range has no Brave equivalent, so it renders to an empty +/// string and is omitted from the query; the operator-visible note is a +/// caller concern. +fn brave_freshness(freshness: &Freshness) -> String { + match freshness { + Freshness::Day => "pd".to_owned(), + Freshness::Week => "pw".to_owned(), + Freshness::Month => "pm".to_owned(), + Freshness::Year => "py".to_owned(), + Freshness::Range { .. } => { + tracing::debug!("web_search freshness date range ignored: Brave only supports pd/pw/pm/py"); + String::new() + } + } +} + +/// Clamps a requested count to Brave's per-section cap without failing the +/// call: models cannot be expected to know the provider's limit. +fn clamp_count(count: u16) -> Result { + let valid = validate_count(count)?; + let clamped = valid.min(BRAVE_MAX_COUNT); + if clamped != valid { + tracing::debug!( + requested = valid, + clamped, + "web_search count clamped to Brave's per-section cap" + ); + } + Ok(clamped) +} + +fn clean_base_url(value: &str) -> Option { + let trimmed = value.trim().trim_end_matches('/'); + (!trimmed.is_empty()).then(|| trimmed.to_owned()) +} + +/// Brave's `GET /res/v1/web/search` response envelope. +/// +/// Forward-compatible: all fields default and unknown keys are tolerated so a +/// Brave API change does not fail the whole search. Web and news result items +/// map straight onto [`WebSearchResult`]; Brave's cosmetic fields +/// (`thumbnail_url`, etc.) and unknown keys are dropped. +#[derive(Debug, Default, Deserialize)] +struct BraveSearchResponse { + #[serde(default, deserialize_with = "null_as_default")] + web: BraveResults, + #[serde(default, deserialize_with = "null_as_default")] + news: BraveResults, + #[serde(default, deserialize_with = "null_as_default")] + query: BraveQuery, +} + +#[derive(Debug, Default, Deserialize)] +struct BraveResults { + #[serde(default, deserialize_with = "null_as_default")] + results: Vec, +} + +#[derive(Debug, Default, Deserialize)] +struct BraveQuery { + #[serde(default)] + search_term: Option, +} + +/// One Brave result item. `serde(default)` keeps every field optional so the +/// provider-neutral mapping degrades gracefully. +#[derive(Debug, Default, Deserialize)] +struct BraveResult { + #[serde(default)] + title: Option, + #[serde(default)] + description: Option, + #[serde(default)] + url: Option, +} + +impl BraveSearchResponse { + fn into_provider_response(self, query: &str, filter: &DomainFilter) -> WebSearchProviderResponse { + let mut web: Vec = self.web.results.into_iter().map(into_web_search_result).collect(); + let mut news: Vec = self.news.results.into_iter().map(into_web_search_result).collect(); + filter.retain(&mut web); + filter.retain(&mut news); + WebSearchProviderResponse { + web, + news, + metadata: WebSearchProviderMetadata { + provider: WebSearchProviderKind::Brave, + query: self.query.search_term.unwrap_or_else(|| query.to_owned()), + search_uuid: None, + latency: None, + }, + } + } +} + +fn into_web_search_result(result: BraveResult) -> WebSearchResult { + WebSearchResult { + url: result.url.unwrap_or_default(), + title: result.title, + description: result.description, + ..WebSearchResult::default() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::types::tools::{WebSearchFilters, WebSearchUserLocation}; + + fn args(json: &str) -> WebSearchArguments { + WebSearchArguments::from_json(json).unwrap() + } + + #[test] + fn api_key_debug_is_redacted() { + let provider = BraveSearchProvider::from_values( + Arc::new(reqwest::Client::new()), + Some("super-secret-key".to_owned()), + Some("https://api.example".to_owned()), + ); + let rendered = format!("{provider:?}"); + assert!(!rendered.contains("super-secret-key")); + assert!(rendered.contains("ApiKey()")); + assert_eq!(format!("{:?}", ApiKey("k".to_owned())), "ApiKey()"); + } + + #[test] + fn from_values_defaults_base_url_and_treats_blank_credentials_as_unset() { + let provider = BraveSearchProvider::from_values(Arc::new(reqwest::Client::new()), None, None); + assert!(provider.api_key.is_none()); + assert_eq!(provider.base_url.as_deref(), Some(BRAVE_DEFAULT_BASE_URL)); + + let provider = BraveSearchProvider::from_values( + Arc::new(reqwest::Client::new()), + Some(" ".to_owned()), + Some(" https://api.example/// ".to_owned()), + ); + assert!(provider.api_key.is_none()); + assert_eq!(provider.base_url.as_deref(), Some("https://api.example")); + + let provider = BraveSearchProvider::from_values(Arc::new(reqwest::Client::new()), Some("k".to_owned()), None); + assert_eq!(provider.base_url.as_deref(), Some(BRAVE_DEFAULT_BASE_URL)); + } + + #[test] + fn provider_caps_concurrency_to_one() { + let provider = BraveSearchProvider::from_values(Arc::new(reqwest::Client::new()), Some("k".to_owned()), None); + assert_eq!(provider.max_concurrent_requests(), Some(BRAVE_MAX_CONCURRENT_REQUESTS)); + } + + #[test] + fn request_renders_brave_query_params() { + let args = args( + r#"{"query":" rust ","count":5,"freshness":"day","country":"us","language":"en","exclude_domains":["a.example"]}"#, + ); + let request = + BraveSearchRequest::from_args_and_config(" rust ", &args, &WebSearchToolParam::default()).unwrap(); + let expected = [ + ("q", "rust"), + ("result_filter", "web,news"), + ("count", "5"), + ("freshness", "pd"), + ("country", "US"), + ("search_lang", "en"), + ] + .map(|(key, value)| (key.to_owned(), value.to_owned())); + assert_eq!(request.query_params(), expected); + } + + #[test] + fn request_clamps_count_to_brave_cap() { + let args = args(r#"{"query":"rust","count":50}"#); + let request = BraveSearchRequest::from_args_and_config("rust", &args, &WebSearchToolParam::default()).unwrap(); + assert_eq!(request.count, Some(BRAVE_MAX_COUNT)); + } + + #[test] + fn request_drops_unrenderable_date_range_freshness() { + let args = args(r#"{"query":"rust","freshness":"2024-01-01to2024-02-01"}"#); + let request = BraveSearchRequest::from_args_and_config("rust", &args, &WebSearchToolParam::default()).unwrap(); + let expected = + [("q", "rust"), ("result_filter", "web,news")].map(|(key, value)| (key.to_owned(), value.to_owned())); + assert_eq!(request.query_params(), expected); + } + + #[test] + fn request_applies_context_size_default_and_tool_config_overrides() { + let config = WebSearchToolParam { + search_context_size: Some(WebSearchContextSize::High), + filters: Some(WebSearchFilters { + allowed_domains: Some(vec![" docs.example ".to_owned()]), + blocked_domains: Some(vec!["bad.example".to_owned()]), + }), + user_location: Some(WebSearchUserLocation { + country: Some(" de ".to_owned()), + ..WebSearchUserLocation::default() + }), + }; + let args = args(r#"{"query":"rust","country":"us","include_domains":["other.example"]}"#); + let request = BraveSearchRequest::from_args_and_config("rust", &args, &config).unwrap(); + // Config filters win over argument filters; allowlist is hard. + assert!(request.domain_filter.allows("https://docs.example/page")); + assert!(!request.domain_filter.allows("https://bad.example/page")); + assert!(!request.domain_filter.allows("https://other.example/page")); + assert_eq!(request.country.as_deref(), Some("DE")); + } + + #[test] + fn domain_filter_respects_label_boundaries() { + let filter = DomainFilter::new(Some(&["example.com".to_owned()]), Some(&["a.example.com".to_owned()])); + assert!(filter.allows("https://example.com/")); + assert!(filter.allows("https://sub.example.com/x")); + assert!(!filter.allows("https://notexample.com/")); + assert!(!filter.allows("https://a.example.com/")); + // Unparseable URL is rejected whenever any filter is active (fail closed). + assert!(!filter.allows("not a url")); + } + + #[test] + fn response_maps_documented_fields_and_tolerates_nulls() { + let response: BraveSearchResponse = serde_json::from_str( + r#"{ + "web": {"results": [ + {"title": "Rust", "description": "desc", "url": "https://example.com/rust", "unknown_key": 1} + ]}, + "news": {"results": []}, + "query": {"search_term": "rust"} + }"#, + ) + .unwrap(); + let mapped = response.into_provider_response("rust", &DomainFilter::default()); + assert!(mapped.news.is_empty()); + assert_eq!(mapped.metadata.provider, WebSearchProviderKind::Brave); + assert_eq!(mapped.metadata.query, "rust"); + let result = &mapped.web[0]; + assert_eq!(result.url, "https://example.com/rust"); + assert_eq!(result.title.as_deref(), Some("Rust")); + assert_eq!(result.description.as_deref(), Some("desc")); + } + + #[test] + fn response_applies_domain_post_filter() { + let response: BraveSearchResponse = serde_json::from_str( + r#"{ + "web": {"results": [ + {"title": "A", "url": "https://example.com/a"}, + {"title": "B", "url": "https://bad.example/b"} + ]}, + "query": {} + }"#, + ) + .unwrap(); + let filter = DomainFilter::new(None, Some(&["bad.example".to_owned()])); + let mapped = response.into_provider_response("rust", &filter); + assert_eq!(mapped.web.len(), 1); + assert_eq!(mapped.web[0].url, "https://example.com/a"); + } + + #[test] + fn response_tolerates_missing_envelope_sections() { + let response: BraveSearchResponse = serde_json::from_str("{}").unwrap(); + let mapped = response.into_provider_response("rust", &DomainFilter::default()); + assert!(mapped.web.is_empty()); + assert!(mapped.news.is_empty()); + assert_eq!(mapped.metadata.query, "rust"); + } +} diff --git a/crates/agentic-server-core/src/tool/web_search/mod.rs b/crates/agentic-server-core/src/tool/web_search/mod.rs index 3f530d60..550fa892 100644 --- a/crates/agentic-server-core/src/tool/web_search/mod.rs +++ b/crates/agentic-server-core/src/tool/web_search/mod.rs @@ -7,6 +7,7 @@ //! as [`you`] shape requests and map responses. pub(crate) mod args; +pub(crate) mod brave; pub(crate) mod you; use std::collections::HashMap; @@ -21,7 +22,8 @@ use serde::{Deserialize, Deserializer, Serialize}; use serde_json::Value; use tokio::sync::Semaphore; -use self::args::{MAX_WEB_SEARCH_QUERIES, WebSearchArguments}; +use self::args::{MAX_WEB_SEARCH_QUERIES, WebSearchArguments, cap_provider_concurrency}; +use self::brave::BraveSearchProvider; use self::you::{YOU_API_BASE_URL, YOU_API_KEY, YouSearchProvider}; use super::handler::MAX_GATEWAY_TOOL_OUTPUT_BYTES; use super::handler::{GatewayExecutor, GatewayToolEventPlan, ToolError, ToolHandler, ToolOutput}; @@ -167,26 +169,30 @@ pub struct WebSearchHandler { impl WebSearchHandler { #[must_use] pub fn from_env(client: Arc) -> Self { - Self::from_values( + Self::from_provider( client, + WebSearchProviderKind::You, std::env::var(YOU_API_KEY).ok(), std::env::var(YOU_API_BASE_URL).ok(), DEFAULT_MAX_CONCURRENT_GATEWAY_CALLS, ) } - /// Builds the You.com-backed handler. - /// - /// `max_concurrent_queries` is the gateway-wide ceiling; the provider's own - /// [`WebSearchProvider::max_concurrent_requests`] ceiling caps it again. + /// Builds the handler for the given provider from optional credential and + /// base URL values; the provider's [`WebSearchProvider::max_concurrent_requests`] + /// ceiling further caps `max_concurrent_queries`. #[must_use] - pub fn from_values( + pub fn from_provider( client: Arc, + provider_kind: WebSearchProviderKind, api_key: Option, base_url: Option, max_concurrent_queries: NonZeroUsize, ) -> Self { - let provider = Arc::new(YouSearchProvider::from_values(client, api_key, base_url)); + let provider: Arc = match provider_kind { + WebSearchProviderKind::You => Arc::new(YouSearchProvider::from_values(client, api_key, base_url)), + WebSearchProviderKind::Brave => Arc::new(BraveSearchProvider::from_values(client, api_key, base_url)), + }; let effective = effective_query_concurrency(provider.as_ref(), max_concurrent_queries); Self::with_provider_and_query_concurrency(provider, effective) } @@ -296,16 +302,11 @@ impl WebSearchHandler { /// Caps the requested query concurrency at the provider's own ceiling. fn effective_query_concurrency(provider: &dyn WebSearchProvider, requested: NonZeroUsize) -> NonZeroUsize { - provider - .max_concurrent_requests() - .map_or(requested, |ceiling| requested.min(ceiling)) + cap_provider_concurrency(provider.max_concurrent_requests(), requested) } -/// A search backend behind `web_search`. -/// -/// Implementations shape one provider request per query and normalize the -/// response into [`WebSearchProviderResponse`]; the handler owns fan-out, -/// concurrency, and the model-facing output shape. +/// A search backend behind `web_search`: implementations shape one request per +/// query and normalize the response into [`WebSearchProviderResponse`]. pub(crate) trait WebSearchProvider: std::fmt::Debug + Send + Sync { fn search<'a>( &'a self, @@ -314,8 +315,7 @@ pub(crate) trait WebSearchProvider: std::fmt::Debug + Send + Sync { config: &'a WebSearchToolParam, ) -> Pin> + Send + 'a>>; - /// Provider-imposed ceiling on concurrent requests, if any. The handler - /// never schedules more queries at once than this allows. + /// Provider-imposed ceiling on concurrent requests, if any. fn max_concurrent_requests(&self) -> Option { None } @@ -802,8 +802,9 @@ mod tests { #[test] fn from_values_inherits_gateway_concurrency() { - let handler = WebSearchHandler::from_values( + let handler = WebSearchHandler::from_provider( Arc::new(reqwest::Client::new()), + WebSearchProviderKind::You, None, None, NonZeroUsize::new(7).expect("nonzero test limit"), @@ -815,8 +816,9 @@ mod tests { #[test] fn from_values_does_not_leak_api_key_in_debug_output() { - let handler = WebSearchHandler::from_values( + let handler = WebSearchHandler::from_provider( Arc::new(reqwest::Client::new()), + WebSearchProviderKind::You, Some("super-secret-key".to_owned()), Some("https://api.example".to_owned()), DEFAULT_MAX_CONCURRENT_GATEWAY_CALLS, diff --git a/crates/agentic-server-core/tests/web_search_tool_test.rs b/crates/agentic-server-core/tests/web_search_tool_test.rs index d54e156d..087b5caa 100644 --- a/crates/agentic-server-core/tests/web_search_tool_test.rs +++ b/crates/agentic-server-core/tests/web_search_tool_test.rs @@ -4,6 +4,7 @@ use std::sync::{ }; use std::time::Duration; +use agentic_core::config::WebSearchProviderKind; use agentic_core::executor::{ConversationHandler, ExecuteRequest, ExecutionContext, ResponseHandler}; use agentic_core::storage::{ConversationStore, ResponseStore}; use agentic_core::tool::{GatewayExecutor, ToolOutput, WebSearchHandler}; @@ -13,7 +14,7 @@ use agentic_core::types::io::{ FunctionToolResultMessage, InputItem, OutputItem, ResponsesInput, ToolCallOutput, ToolChoice, }; use agentic_core::types::request_response::{RequestPayload, ResponseTextConfig}; -use agentic_core::types::tools::{ResponsesTool, WebSearchToolParam}; +use agentic_core::types::tools::{ResponsesTool, WebSearchFilters, WebSearchToolParam}; use axum::extract::State; use axum::http::{HeaderMap, StatusCode, Uri}; use axum::routing::{get, post}; @@ -2540,3 +2541,320 @@ async fn stream_returns_incomplete_after_max_gateway_tool_rounds() { } assert_eq!(llm.request_bodies().await.len(), 10); } + +// --------------------------------------------------------------------------- +// Brave provider +// --------------------------------------------------------------------------- + +#[derive(Debug)] +struct CapturedBraveRequest { + token: String, + body: serde_json::Value, +} + +/// Aborts the mock server task when dropped, so no test leaves a runtime +/// worker behind. +struct MockBraveServer { + handle: tokio::task::JoinHandle<()>, +} + +impl Drop for MockBraveServer { + fn drop(&mut self) { + self.handle.abort(); + } +} + +/// Mock of Brave's `GET /res/v1/web/search` envelope. Returns `retry_after` in +/// the header on failure so the 429 propagation path can be asserted. +async fn spawn_mock_brave( + status: StatusCode, + response_body: serde_json::Value, + retry_after: Option<&str>, +) -> (String, mpsc::Receiver, MockBraveServer) { + let retry_after = retry_after.map(str::to_owned); + let (tx, rx) = mpsc::channel(16); + let app = Router::new() + .route( + "/res/v1/web/search", + get( + move |State(tx): State>, headers: HeaderMap, uri: Uri| async move { + let token = headers + .get("x-subscription-token") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_owned(); + let body = query_params_as_json(&uri); + tx.send(CapturedBraveRequest { token, body }).await.unwrap(); + let json = serde_json::to_string(&response_body).unwrap(); + let mut response = axum::response::Response::new(axum::body::Body::from(json)); + *response.status_mut() = status; + response + .headers_mut() + .insert(axum::http::header::CONTENT_TYPE, "application/json".parse().unwrap()); + if let Some(retry) = retry_after { + response + .headers_mut() + .insert(axum::http::header::RETRY_AFTER, retry.parse().unwrap()); + } + response + }, + ), + ) + .with_state(tx); + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let handle = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + (format!("http://{addr}"), rx, MockBraveServer { handle }) +} + +#[tokio::test] +async fn brave_handler_executes_against_mock_brave_endpoint() { + let (base_url, mut captured, _handle) = spawn_mock_brave( + StatusCode::OK, + serde_json::json!({ + "web": {"results": [ + {"title": "Rust", "description": "a", "url": "https://rust.example/rust"} + ]}, + "news": {"results": [ + {"title": "News", "description": "b", "url": "https://news.example/x"} + ]}, + "query": {"search_term": "rust async"} + }), + None, + ) + .await; + let handler = WebSearchHandler::from_provider( + Arc::new(reqwest::Client::new()), + WebSearchProviderKind::Brave, + Some("secret-brave-key".to_owned()), + Some(base_url), + std::num::NonZeroUsize::new(5).unwrap(), + ); + let output = handler + .execute( + "call_brave", + "web_search", + r#"{"query":"rust async","count":50}"#, + &WebSearchToolParam::default(), + ) + .await + .expect("brave search should succeed"); + let request = captured.recv().await.expect("mock Brave should receive the request"); + assert_eq!(request.token, "secret-brave-key"); + // count 50 clamped to Brave's cap of 20. + assert_eq!(request.body["count"], 20); + assert_eq!(request.body["q"], "rust async"); + assert_eq!(request.body["result_filter"], "web,news"); + let parsed: serde_json::Value = serde_json::from_str(&output.output).expect("brave output JSON"); + assert_eq!(parsed["results"]["web"][0]["url"], "https://rust.example/rust"); + assert_eq!(parsed["results"]["news"][0]["url"], "https://news.example/x"); + assert_eq!(parsed["metadata"][0]["query"], "rust async"); +} + +#[tokio::test] +async fn brave_handler_surfaces_401_and_429_with_retry_after() { + let (base_url, _captured, _handle) = spawn_mock_brave( + StatusCode::UNAUTHORIZED, + serde_json::json!({"error":"unauthorized"}), + None, + ) + .await; + let handler = WebSearchHandler::from_provider( + Arc::new(reqwest::Client::new()), + WebSearchProviderKind::Brave, + Some("secret-brave-key".to_owned()), + Some(base_url), + std::num::NonZeroUsize::new(1).unwrap(), + ); + let err = handler + .execute( + "call_brave", + "web_search", + r#"{"query":"rust"}"#, + &WebSearchToolParam::default(), + ) + .await + .expect_err("401 should fail the call"); + let message = format!("{err}"); + assert!(message.contains("401"), "should surface the status: {message}"); + assert!( + !message.contains("secret-brave-key"), + "must not leak the secret: {message}" + ); + + let (base_url, _captured, _handle) = + spawn_mock_brave(StatusCode::TOO_MANY_REQUESTS, serde_json::json!({}), Some("30")).await; + let handler = WebSearchHandler::from_provider( + Arc::new(reqwest::Client::new()), + WebSearchProviderKind::Brave, + Some("secret-brave-key".to_owned()), + Some(base_url), + std::num::NonZeroUsize::new(1).unwrap(), + ); + let err = handler + .execute( + "call_brave", + "web_search", + r#"{"query":"rust"}"#, + &WebSearchToolParam::default(), + ) + .await + .expect_err("429 should fail the call"); + let message = format!("{err}"); + assert!(message.contains("429"), "should surface the status: {message}"); + assert!( + message.contains("retry after 30"), + "should surface Retry-After verbatim: {message}" + ); +} + +#[tokio::test] +async fn brave_handler_domain_post_filter_drops_blocklisted_hosts() { + let (base_url, _captured, _handle) = spawn_mock_brave( + StatusCode::OK, + serde_json::json!({ + "web": {"results": [ + {"title": "A", "url": "https://docs.example/rust"}, + {"title": "B", "url": "https://bad.example/rust"} + ]}, + "query": {"search_term": "rust"} + }), + None, + ) + .await; + let handler = WebSearchHandler::from_provider( + Arc::new(reqwest::Client::new()), + WebSearchProviderKind::Brave, + Some("secret-brave-key".to_owned()), + Some(base_url), + std::num::NonZeroUsize::new(1).unwrap(), + ); + let config = WebSearchToolParam { + filters: Some(WebSearchFilters { + blocked_domains: Some(vec!["bad.example".to_owned()]), + ..Default::default() + }), + ..WebSearchToolParam::default() + }; + let output = handler + .execute("call_brave", "web_search", r#"{"query":"rust"}"#, &config) + .await + .expect("search should succeed"); + let parsed: serde_json::Value = serde_json::from_str(&output.output).expect("output JSON"); + let urls: Vec<&str> = parsed["results"]["web"] + .as_array() + .unwrap() + .iter() + .map(|r| r["url"].as_str().unwrap()) + .collect(); + assert_eq!( + urls, + vec!["https://docs.example/rust"], + "blocklisted host must be dropped" + ); +} + +#[tokio::test] +async fn brave_handler_empty_results_are_graceful() { + let (base_url, _captured, _handle) = spawn_mock_brave( + StatusCode::OK, + serde_json::json!({"web":{"results":[]},"query":{"search_term":"nothing"}}), + None, + ) + .await; + let handler = WebSearchHandler::from_provider( + Arc::new(reqwest::Client::new()), + WebSearchProviderKind::Brave, + Some("secret-brave-key".to_owned()), + Some(base_url), + std::num::NonZeroUsize::new(1).unwrap(), + ); + let output = handler + .execute( + "call_brave", + "web_search", + r#"{"query":"nothing"}"#, + &WebSearchToolParam::default(), + ) + .await + .expect("empty results should not fail"); + let parsed: serde_json::Value = serde_json::from_str(&output.output).expect("output JSON"); + assert_eq!(parsed["results"]["web"].as_array().unwrap().len(), 0); +} + +/// The free-tier ~1 QPS contract: whatever the gateway-wide limit requests, a +/// Brave-backed handler never has more than one query request in flight, and +/// every query still reaches the backend exactly once. The mock server counts +/// concurrent in-flight requests plus total requests, and the batched call +/// must complete only when the three queries were serialized. +#[tokio::test] +async fn brave_provider_caps_effective_query_concurrency_to_one() { + use std::sync::atomic::{AtomicUsize, Ordering}; + + #[derive(Clone, Default)] + struct Counters { + current: Arc, + max_seen: Arc, + total: Arc, + } + + let counters = Counters::default(); + let max_seen = std::sync::Arc::clone(&counters.max_seen); + let total = std::sync::Arc::clone(&counters.total); + let app = Router::new().route( + "/res/v1/web/search", + get(move || { + let counters = counters.clone(); + async move { + let active = counters.current.fetch_add(1, Ordering::SeqCst) + 1; + counters.max_seen.fetch_max(active, Ordering::SeqCst); + counters.total.fetch_add(1, Ordering::SeqCst); + // Hold the request long enough that an unscheduled second query + // would overlap with this one. + tokio::time::sleep(Duration::from_millis(50)).await; + counters.current.fetch_sub(1, Ordering::SeqCst); + axum::Json(serde_json::json!({"web": {"results": []}})) + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + // RAII guard so a panic between spawn and assertions cannot leak the task. + // Kept for its Drop impl; the leading underscore silences the unused warning + // while still holding the guard alive until the end of the test. + let _server = MockBraveServer { + handle: tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }), + }; + + let handler = WebSearchHandler::from_provider( + Arc::new(reqwest::Client::new()), + WebSearchProviderKind::Brave, + Some("secret-brave-key".to_owned()), + Some(format!("http://{addr}")), + // Gateway-wide ceiling deliberately above the Brave provider's own + // ceiling of 1: the effective limit must still be 1. + std::num::NonZeroUsize::new(5).unwrap(), + ); + let output = handler + .execute( + "call_brave", + "web_search", + r#"{"queries":["a","b","c"]}"#, + &WebSearchToolParam::default(), + ) + .await + .expect("batched brave search should succeed"); + let parsed: serde_json::Value = serde_json::from_str(&output.output).expect("output JSON"); + assert_eq!(parsed["metadata"].as_array().unwrap().len(), 3); + assert_eq!( + total.load(Ordering::SeqCst), + 3, + "every query must reach the backend exactly once (no silent drops)" + ); + assert_eq!( + max_seen.load(Ordering::SeqCst), + 1, + "Brave requests must be serialized even with a higher gateway-wide limit" + ); +} diff --git a/crates/agentic-server/src/config_file.rs b/crates/agentic-server/src/config_file.rs index c260f6af..c710cf4c 100644 --- a/crates/agentic-server/src/config_file.rs +++ b/crates/agentic-server/src/config_file.rs @@ -4,7 +4,7 @@ use std::num::NonZeroUsize; use std::path::Path; use agentic_core::McpServerEntry; -use agentic_core::config::CONFIG_FILE_NAME; +use agentic_core::config::{CONFIG_FILE_NAME, WebSearchProviderKind}; use agentic_core::error::Error; use serde::{Deserialize, Serialize}; @@ -15,11 +15,13 @@ pub(crate) struct WebSearchFileConfig { pub base_url: Option, #[serde(skip_serializing_if = "Option::is_none")] pub api_key_env: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub provider: Option, } impl WebSearchFileConfig { fn is_empty(&self) -> bool { - self.base_url.is_none() && self.api_key_env.is_none() + self.base_url.is_none() && self.api_key_env.is_none() && self.provider.is_none() } } @@ -269,6 +271,7 @@ mod tests { web_search: WebSearchFileConfig { base_url: Some("https://api.ydc-index.io".to_owned()), api_key_env: Some("YOU_API_KEY".to_owned()), + provider: None, }, mcp: McpFileConfig { allowed_hosts: vec!["mcp.example.com".to_owned()], diff --git a/crates/agentic-server/src/main.rs b/crates/agentic-server/src/main.rs index 008845e6..a81e31b3 100644 --- a/crates/agentic-server/src/main.rs +++ b/crates/agentic-server/src/main.rs @@ -11,7 +11,8 @@ use agentic_core::config::{ DEFAULT_POSTGRES_MAX_LIFETIME_SECONDS, DEFAULT_POSTGRES_MIGRATION_TIMEOUT_SECONDS, DEFAULT_POSTGRES_STATEMENT_TIMEOUT_SECONDS, DEFAULT_SQLITE_JOURNAL_SIZE_LIMIT_BYTES, DEFAULT_SQLITE_MAX_CONNECTIONS, DEFAULT_SQLITE_MMAP_SIZE_BYTES, PostgresConfig, SqliteConfig, SqliteTempStore, - ToolRuntimeConfig, WebSearchProviderConfig, default_database_url, ensure_agentic_api_home, normalize_base_url, + ToolRuntimeConfig, WebSearchProviderConfig, WebSearchProviderKind, default_database_url, ensure_agentic_api_home, + normalize_base_url, }; use agentic_core::error::Error; use agentic_server::app::DEFAULT_MAX_REQUEST_BODY_SIZE; @@ -303,8 +304,18 @@ fn build_config(llm_api_base: String, common: &CommonArgs, file: &FileConfig) -> .or_else(|| file.database_url.clone()) .map_or_else(default_database_url, Ok)?; let (postgres, sqlite) = database_configs_from_env(&db_url)?; - let web_search_api_key = file.web_search.api_key_env.as_deref().and_then(environment_value); - let web_search_base_url = environment_value("YOU_API_BASE_URL").or_else(|| file.web_search.base_url.clone()); + let web_search_provider_kind = web_search_provider_kind(file)?; + // Fall back to the selected provider's default credential variable so + // switching providers never requires editing the generated config file. + let web_search_api_key_env = web_search_api_key_env(file.web_search.api_key_env.clone(), web_search_provider_kind); + let web_search_api_key = environment_value(&web_search_api_key_env); + let web_search_base_url = match web_search_provider_kind { + // You.com keeps its historical override name; other providers + // (including Brave) use the generic `AGENTIC_WEB_SEARCH_BASE_URL`. + WebSearchProviderKind::You => environment_value("YOU_API_BASE_URL"), + _ => environment_value("AGENTIC_WEB_SEARCH_BASE_URL"), + } + .or_else(|| file.web_search.base_url.clone()); let mcp_allowed_hosts = environment_value("AGENTIC_MCP_ALLOWED_HOSTS") .map_or_else(|| file.mcp.allowed_hosts.clone(), |value| parse_comma_separated(&value)); let max_concurrent_gateway_calls_default = file @@ -325,7 +336,7 @@ fn build_config(llm_api_base: String, common: &CommonArgs, file: &FileConfig) -> postgres, sqlite, tools: ToolRuntimeConfig { - web_search: WebSearchProviderConfig::new(web_search_api_key, web_search_base_url), + web_search: WebSearchProviderConfig::new(web_search_provider_kind, web_search_api_key, web_search_base_url), mcp_servers: file.mcp_servers.clone(), mcp_allowed_hosts, messages_gateway_tool_aliases: file.messages_gateway.tool_aliases.clone(), @@ -334,6 +345,33 @@ fn build_config(llm_api_base: String, common: &CommonArgs, file: &FileConfig) -> }) } +/// Resolves the `web_search` provider: the `AGENTIC_WEB_SEARCH_PROVIDER` +/// environment value overrides the configuration file's `provider` key; +/// both default to You.com. +fn web_search_provider_kind(file: &FileConfig) -> Result { + resolve_web_search_provider_kind( + environment_value("AGENTIC_WEB_SEARCH_PROVIDER"), + file.web_search.provider, + ) +} + +/// The environment variable holding the web-search credential: the file's +/// `api_key_env` wins, otherwise the selected provider's default name. +fn web_search_api_key_env(file_value: Option, kind: WebSearchProviderKind) -> String { + file_value.unwrap_or_else(|| kind.default_api_key_env().to_owned()) +} + +fn resolve_web_search_provider_kind( + environment: Option, + file_value: Option, +) -> Result { + let parse = |value: &str| WebSearchProviderKind::parse_name("AGENTIC_WEB_SEARCH_PROVIDER", value); + match environment { + Some(value) => parse(&value), + None => Ok(file_value.unwrap_or_default()), + } +} + fn gateway_options<'a>( common: &'a CommonArgs, file: &FileConfig, @@ -355,7 +393,8 @@ fn generated_file_config(llm_api_base: String) -> FileConfig { llm_api_base: Some(llm_api_base), web_search: WebSearchFileConfig { base_url: environment_value("YOU_API_BASE_URL"), - api_key_env: Some("YOU_API_KEY".to_owned()), + api_key_env: None, + provider: None, }, mcp: McpFileConfig { allowed_hosts: environment_value("AGENTIC_MCP_ALLOWED_HOSTS") @@ -466,11 +505,14 @@ mod tests { use clap::{CommandFactory, Parser}; + use super::config_file::{FileConfig, WebSearchFileConfig}; use super::{ Cli, Commands, database_configs_from_env, oidc_config_from_values, parse_env_duration_value, parse_env_nonzero_usize_value, parse_env_optional_duration_value, parse_env_temp_store_value, parse_env_u32_value, parse_env_u64_value, resolve_max_request_body_size_value, + resolve_web_search_provider_kind, web_search_api_key_env, web_search_provider_kind, }; + use agentic_core::config::WebSearchProviderKind; use agentic_core::config::{ DEFAULT_POSTGRES_ACQUIRE_TIMEOUT_SECONDS, DEFAULT_POSTGRES_IDLE_TIMEOUT_SECONDS, DEFAULT_POSTGRES_LOCK_TIMEOUT_SECONDS, DEFAULT_POSTGRES_MAX_CONNECTIONS, @@ -831,4 +873,78 @@ mod tests { let error = database_configs_from_env("not a database URL").expect_err("invalid URL must be rejected"); assert!(error.to_string().contains("invalid DATABASE_URL")); } + + #[test] + fn web_search_api_key_env_falls_back_to_provider_default() { + // The file's explicit name wins over the provider default. + assert_eq!( + web_search_api_key_env(Some("CUSTOM_KEY".to_owned()), WebSearchProviderKind::Brave), + "CUSTOM_KEY" + ); + // Without a file entry the selected provider's default env name is used, + // so `AGENTIC_WEB_SEARCH_PROVIDER=brave` works with only `BRAVE_API_KEY` + // set and no hand-edited configuration file. + assert_eq!( + web_search_api_key_env(None, WebSearchProviderKind::Brave), + "BRAVE_API_KEY" + ); + assert_eq!(web_search_api_key_env(None, WebSearchProviderKind::You), "YOU_API_KEY"); + } + + #[test] + fn web_search_provider_defaults_to_you_from_file_config() { + let file = FileConfig { + web_search: WebSearchFileConfig::default(), + ..FileConfig::default() + }; + assert_eq!( + web_search_provider_kind(&file).expect("default provider"), + WebSearchProviderKind::You + ); + + let file = FileConfig { + web_search: WebSearchFileConfig { + provider: Some(WebSearchProviderKind::Brave), + ..WebSearchFileConfig::default() + }, + ..FileConfig::default() + }; + assert_eq!( + web_search_provider_kind(&file).expect("file provider"), + WebSearchProviderKind::Brave + ); + } + + #[test] + fn web_search_provider_environment_value_overrides_file_config() { + // The environment value wins over the file's provider even when the two + // disagree. + assert_eq!( + resolve_web_search_provider_kind(Some("brave".to_owned()), Some(WebSearchProviderKind::You)) + .expect("env override"), + WebSearchProviderKind::Brave + ); + // Without an environment value the file's provider is used. + assert_eq!( + resolve_web_search_provider_kind(None, Some(WebSearchProviderKind::Brave)).expect("file provider"), + WebSearchProviderKind::Brave + ); + // Neither source set: default to You. + assert_eq!( + resolve_web_search_provider_kind(None, None).expect("default"), + WebSearchProviderKind::You + ); + } + + #[test] + fn web_search_provider_rejects_unknown_names() { + let error = resolve_web_search_provider_kind(Some("tavily".to_owned()), None) + .expect_err("unknown provider must be rejected"); + assert!(error.to_string().contains("invalid AGENTIC_WEB_SEARCH_PROVIDER")); + // Whitespace is trimmed, so ` brave ` is still valid. + assert_eq!( + resolve_web_search_provider_kind(Some(" brave ".to_owned()), None).expect("trimmed valid value"), + WebSearchProviderKind::Brave + ); + } } diff --git a/docs/deploying/kubernetes.md b/docs/deploying/kubernetes.md index 6ebcd2b3..5fce0939 100644 --- a/docs/deploying/kubernetes.md +++ b/docs/deploying/kubernetes.md @@ -460,6 +460,32 @@ patches: executed by the gateway. Responses API clients declare web search structurally and do not need this alias. Leaving the setting empty preserves the default client-owned behavior. +#### Switching the web-search provider to Brave + +The gateway supports Brave Search as an alternative `web_search` backend. Set the provider through the ConfigMap and +put its credential in the same Secret: + +```yaml +patches: + - target: + kind: ConfigMap + name: agentic-api + patch: |- + - op: add + path: /data/AGENTIC_WEB_SEARCH_PROVIDER + value: brave + # Optional: overrides the Brave API base URL (default + # https://api.search.brave.com). + - op: add + path: /data/AGENTIC_WEB_SEARCH_BASE_URL + value: https://api.search.brave.com +``` + +The Secret file then contains `BRAVE_API_KEY=...` instead of (or in addition to) `YOU_API_KEY=...`; the provider's +default credential variable is selected automatically, so no `api_key_env` file setting is needed. Brave clamps +`count` to 20 results per section, post-filters domain allow/block lists client-side, and runs at most one request in +flight to respect the free-tier ~1 QPS rate limit. + Apply the overlay, restart the Deployment after every Secret update, and wait for readiness: ```console