pub mod exa;
pub mod serper;
pub mod tavily;
pub use exa::ExaBackend;
pub use serper::SerperBackend;
pub use tavily::TavilyBackend;
use async_trait::async_trait;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use lc_core::tools::{BaseTool, ToolError};
pub const DEFAULT_TOP_K: usize = 5;
pub const MAX_TOP_K: usize = 20;
#[derive(Debug, Clone, Serialize, PartialEq)]
pub struct SearchResult {
pub title: String,
pub url: String,
pub snippet: String,
pub score: f64,
pub published_date: Option<String>,
pub author: Option<String>,
pub provider: &'static str,
}
#[derive(Debug, Clone, Serialize)]
pub struct SearchOutput {
pub query: String,
pub results: Vec<SearchResult>,
pub answer: Option<String>,
pub provider: &'static str,
}
#[derive(Debug, Clone, Default)]
pub struct BackendResponse {
pub results: Vec<SearchResult>,
pub answer: Option<String>,
}
#[async_trait]
pub trait SearchBackend: Send + Sync {
fn label(&self) -> &'static str;
async fn search(
&self,
query: &str,
top_k: usize,
include_answer: bool,
) -> Result<BackendResponse, ToolError>;
}
#[derive(Debug, Deserialize, JsonSchema)]
pub struct HostedSearchInput {
pub query: String,
pub top_k: Option<usize>,
pub include_answer: Option<bool>,
}
pub struct HostedSearchTool {
backend: Arc<dyn SearchBackend>,
tool_name: String,
}
impl HostedSearchTool {
pub fn new(backend: impl SearchBackend + 'static) -> Self {
Self::from_arc(Arc::new(backend))
}
pub fn from_arc(backend: Arc<dyn SearchBackend>) -> Self {
let tool_name = format!("{}_search", backend.label());
Self { backend, tool_name }
}
pub fn provider(&self) -> &'static str {
self.backend.label()
}
pub fn tavily(api_key: impl Into<String>) -> Self {
Self::new(TavilyBackend::new(api_key))
}
pub fn tavily_from_env() -> Result<Self, ToolError> {
Ok(Self::new(TavilyBackend::from_env()?))
}
pub fn serper(api_key: impl Into<String>) -> Self {
Self::new(SerperBackend::new(api_key))
}
pub fn serper_from_env() -> Result<Self, ToolError> {
Ok(Self::new(SerperBackend::from_env()?))
}
pub fn exa(api_key: impl Into<String>) -> Self {
Self::new(ExaBackend::new(api_key))
}
pub fn exa_from_env() -> Result<Self, ToolError> {
Ok(Self::new(ExaBackend::from_env()?))
}
pub async fn search(
&self,
query: &str,
top_k: Option<usize>,
include_answer: Option<bool>,
) -> Result<SearchOutput, ToolError> {
let query = query.trim();
if query.is_empty() {
return Err(ToolError::InvalidInput(
"search query must not be empty".to_string(),
));
}
let top_k = top_k.unwrap_or(DEFAULT_TOP_K).clamp(1, MAX_TOP_K);
let include_answer = include_answer.unwrap_or(true);
let mut response = self.backend.search(query, top_k, include_answer).await?;
rank_results(&mut response.results);
response.results.truncate(top_k);
Ok(SearchOutput {
query: query.to_string(),
results: response.results,
answer: response.answer,
provider: self.backend.label(),
})
}
}
pub(crate) fn canonical_url(raw: &str) -> String {
let raw = raw.trim();
match url::Url::parse(raw) {
Ok(mut parsed) => {
parsed.set_fragment(None);
if let Some(host) = parsed.host_str().map(str::to_lowercase) {
let _ = parsed.set_host(Some(&host));
}
let path = parsed.path().to_string();
if path.len() > 1 && path.ends_with('/') {
parsed.set_path(path.trim_end_matches('/'));
}
parsed.to_string()
}
Err(_) => raw.to_lowercase(),
}
}
pub(crate) fn rank_results(results: &mut Vec<SearchResult>) {
results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
let mut seen = std::collections::HashSet::new();
results.retain(|r| seen.insert(canonical_url(&r.url)));
}
pub(crate) fn require_env(key: &str) -> Result<String, ToolError> {
let value = std::env::var(key).map_err(|_| {
ToolError::InvalidInput(format!("{key} environment variable not set or empty"))
})?;
let value = value.trim();
if value.is_empty() {
return Err(ToolError::InvalidInput(format!(
"{key} environment variable not set or empty"
)));
}
Ok(value.to_string())
}
pub(crate) fn trim_trailing_slash(mut base: String) -> String {
while base.len() > 1 && base.ends_with('/') {
base.pop();
}
base
}
fn render_text(output: &SearchOutput) -> String {
let mut text = format!("{} 搜索结果(查询: {})\n\n", output.provider, output.query);
if let Some(answer) = output.answer.as_ref().filter(|a| !a.is_empty()) {
text.push_str(&format!("综合答案: {answer}\n\n"));
}
for (i, result) in output.results.iter().enumerate() {
text.push_str(&format!("{}. {}\n", i + 1, result.title));
text.push_str(&format!(" {}\n", result.snippet));
text.push_str(&format!(" URL: {}\n", result.url));
text.push_str(&format!(" 相关度: {:.2}\n", result.score));
match (result.published_date.as_ref(), result.author.as_ref()) {
(Some(date), Some(author)) => {
text.push_str(&format!(" 发布: {date} · {author}\n"));
}
(Some(date), None) => text.push_str(&format!(" 发布: {date}\n")),
(None, Some(author)) => text.push_str(&format!(" 作者: {author}\n")),
(None, None) => {}
}
text.push('\n');
}
if output.results.is_empty() {
text.push_str("未找到相关结果");
} else {
text.push_str(&format!("共 {} 条结果", output.results.len()));
}
text
}
#[async_trait]
impl BaseTool for HostedSearchTool {
fn name(&self) -> &str {
&self.tool_name
}
fn description(&self) -> &str {
match self.backend.label() {
"tavily" => "Tavily 托管网页搜索工具(为 RAG/agent 优化的正文片段,可选综合答案)。\n\n参数:\n- query: 搜索关键词\n- top_k: 返回结果数量(默认 5,上限 20)\n- include_answer: 是否请求综合答案(默认 true)\n\n需设置 TAVILY_API_KEY。\n示例: {\"query\": \"Rust 1.85 async closure\", \"top_k\": 5}",
"serper" => "Serper 托管网页搜索工具(Google 结果页)。\n\n参数:\n- query: 搜索关键词\n- top_k: 返回结果数量(默认 5,上限 20)\n- include_answer: 此后端不支持综合答案,字段被忽略\n\n需设置 SERPER_API_KEY。\n示例: {\"query\": \"tokio tungstenite connect_async\", \"top_k\": 5}",
"exa" => "Exa 神经托管网页搜索工具(语义检索,适合研究类查询)。\n\n参数:\n- query: 搜索查询(自然语言描述)\n- top_k: 返回结果数量(默认 5,上限 20)\n- include_answer: 此后端不支持综合答案,字段被忽略\n\n需设置 EXA_API_KEY。\n示例: {\"query\": \"papers about long-context transformer memory\", \"top_k\": 5}",
_ => "Hosted web search tool.\n\n参数:\n- query: 搜索关键词\n- top_k: 返回结果数量(默认 5,上限 20)",
}
}
async fn run(&self, input: String) -> Result<String, ToolError> {
let parsed: HostedSearchInput = serde_json::from_str(&input)
.map_err(|e| ToolError::InvalidInput(format!("JSON parse failed: {e}")))?;
let output = self
.search(&parsed.query, parsed.top_k, parsed.include_answer)
.await?;
Ok(render_text(&output))
}
fn args_schema(&self) -> Option<serde_json::Value> {
serde_json::to_value(schemars::schema_for!(HostedSearchInput)).ok()
}
}
#[cfg(test)]
pub(crate) mod test_support {
use std::sync::Mutex;
pub(crate) static ENV_LOCK: Mutex<()> = Mutex::new(());
pub(crate) async fn spawn_one_shot_json(
reply: serde_json::Value,
) -> (String, tokio::sync::oneshot::Receiver<Vec<u8>>) {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (tx, rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = Vec::new();
let mut buf = [0u8; 4096];
loop {
let n = socket.read(&mut buf).await.unwrap();
assert!(n > 0, "client closed request early");
request.extend_from_slice(&buf[..n]);
if request.windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let header_end = request
.windows(4)
.position(|w| w == b"\r\n\r\n")
.map(|p| p + 4)
.unwrap();
let content_length = String::from_utf8_lossy(&request[..header_end])
.lines()
.find_map(|line| {
let line = line.to_ascii_lowercase();
line.strip_prefix("content-length:")
.map(|v| v.trim().parse::<usize>().unwrap())
})
.unwrap_or(0);
while request.len() < header_end + content_length {
let n = socket.read(&mut buf).await.unwrap();
assert!(n > 0);
request.extend_from_slice(&buf[..n]);
}
let body = serde_json::to_vec(&reply).unwrap();
let response = format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
socket.write_all(response.as_bytes()).await.unwrap();
socket.write_all(&body).await.unwrap();
socket.flush().await.unwrap();
let _ = tx.send(request);
});
(format!("http://{addr}"), rx)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn backend_debug_never_exposes_api_keys() {
let secret = "supersecret-key-DEBUG-LEAK";
for rendered in [
format!("{:?}", TavilyBackend::new(secret)),
format!("{:?}", SerperBackend::new(secret)),
format!("{:?}", ExaBackend::new(secret)),
] {
assert!(
!rendered.contains(secret),
"key leaked via Debug: {rendered}"
);
assert!(rendered.contains("<redacted>"), "{rendered}");
}
}
#[test]
fn canonical_url_strips_fragment_and_trailing_slash() {
assert_eq!(
canonical_url("HTTPS://Example.COM/path/#section"),
"https://example.com/path"
);
assert_eq!(
canonical_url("https://a.example/x?b=1#frag"),
"https://a.example/x?b=1"
);
assert_eq!(canonical_url("http://b.example/"), "http://b.example/");
assert_eq!(canonical_url(" NOT-A-URL "), "not-a-url");
}
#[test]
fn rank_orders_by_score_and_dedupes_url() {
let provider = "test";
let mk = |url: &str, score: f64| SearchResult {
title: url.to_string(),
url: url.to_string(),
snippet: String::new(),
score,
published_date: None,
author: None,
provider,
};
let mut results = vec![
mk("https://a.example/p", 0.2),
mk("https://a.example/p#frag", 0.9), mk("https://b.example/", 0.5),
];
rank_results(&mut results);
assert_eq!(results.len(), 2, "fragment dup must collapse");
assert_eq!(results[0].url, "https://a.example/p#frag");
assert!(results[0].score > results[1].score);
}
struct RecordingBackend {
seen_k: tokio::sync::Mutex<Option<usize>>,
}
#[async_trait]
impl SearchBackend for RecordingBackend {
fn label(&self) -> &'static str {
"recording"
}
async fn search(
&self,
_query: &str,
top_k: usize,
_include_answer: bool,
) -> Result<BackendResponse, ToolError> {
*self.seen_k.lock().await = Some(top_k);
Ok(BackendResponse::default())
}
}
fn block_on<F: std::future::Future>(fut: F) -> F::Output {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(fut)
}
#[test]
fn empty_query_is_rejected_before_backend() {
let tool = HostedSearchTool::new(RecordingBackend {
seen_k: tokio::sync::Mutex::new(None),
});
let err = block_on(tool.search(" ", None, None)).unwrap_err();
assert!(err.to_string().contains("empty"));
}
#[test]
fn top_k_is_clamped_before_dispatch() {
let backend = Arc::new(RecordingBackend {
seen_k: tokio::sync::Mutex::new(None),
});
let tool = HostedSearchTool::from_arc(backend.clone() as Arc<dyn SearchBackend>);
block_on(tool.search("q", Some(999), None)).unwrap();
assert_eq!(block_on(backend.seen_k.lock()).as_ref(), Some(&MAX_TOP_K));
let backend2 = Arc::new(RecordingBackend {
seen_k: tokio::sync::Mutex::new(None),
});
let tool2 = HostedSearchTool::from_arc(backend2.clone() as Arc<dyn SearchBackend>);
block_on(tool2.search("q", Some(0), None)).unwrap();
assert_eq!(block_on(backend2.seen_k.lock()).as_ref(), Some(&1));
}
#[test]
fn tool_metadata_is_provider_named() {
let tool = HostedSearchTool::tavily("tvly-test");
assert_eq!(BaseTool::name(&tool), "tavily_search");
assert!(tool.description().contains("Tavily"));
assert!(BaseTool::args_schema(&tool).is_some());
assert_eq!(tool.provider(), "tavily");
}
#[test]
fn render_includes_answer_scores_and_empty_state() {
let output = SearchOutput {
query: "q".into(),
results: vec![SearchResult {
title: "T".into(),
url: "https://x.example".into(),
snippet: "S".into(),
score: 0.91,
published_date: Some("2026-01-02".into()),
author: Some("A".into()),
provider: "tavily",
}],
answer: Some("synth".into()),
provider: "tavily",
};
let text = render_text(&output);
assert!(text.contains("综合答案: synth"));
assert!(text.contains("相关度: 0.91"));
assert!(text.contains("2026-01-02 · A"));
let empty = SearchOutput {
query: "q".into(),
results: vec![],
answer: None,
provider: "serper",
};
assert!(render_text(&empty).contains("未找到相关结果"));
}
}