1use super::{SearchProviderConfig, SearchResult, SearchResults};
2use anyhow::Result;
3use serde::{Deserialize, Serialize};
4
5#[derive(Debug, Serialize)]
6struct ExaSearchRequest {
7 query: String,
8 #[serde(skip_serializing_if = "Option::is_none")]
9 num_results: Option<usize>,
10 contents: ExaContentsRequest,
11}
12
13#[derive(Debug, Serialize)]
14struct ExaContentsRequest {
15 text: bool,
16}
17
18#[derive(Debug, Deserialize)]
19struct ExaSearchResponse {
20 results: Vec<ExaResult>,
21}
22
23#[derive(Debug, Deserialize)]
24struct ExaResult {
25 title: String,
26 url: String,
27 #[serde(default)]
28 text: Option<String>,
29 #[serde(rename = "publishedDate")]
30 published_date: Option<String>,
31 author: Option<String>,
32 score: Option<f64>,
33}
34
35pub async fn search(
36 provider_config: &SearchProviderConfig,
37 query: &str,
38 count: Option<usize>,
39) -> Result<SearchResults> {
40 let client = reqwest::Client::builder()
41 .timeout(std::time::Duration::from_secs(30))
42 .build()?;
43
44 let request_body = ExaSearchRequest {
45 query: query.to_string(),
46 num_results: count,
47 contents: ExaContentsRequest { text: true },
48 };
49
50 let base_url = provider_config.url.trim_end_matches('/');
53 let url = if base_url.ends_with("/search") {
54 base_url.to_string()
56 } else {
57 format!("{}/search", base_url)
59 };
60 let mut request = client.post(&url).json(&request_body);
61
62 for (name, value) in &provider_config.headers {
64 request = request.header(name, value);
65 }
66
67 let start_time = std::time::Instant::now();
68 let response = request.send().await?;
69 let search_time_ms = start_time.elapsed().as_millis() as u64;
70
71 if !response.status().is_success() {
72 let status = response.status();
73 let error_text = response.text().await.unwrap_or_default();
74 anyhow::bail!("Exa search API error ({}): {}", status, error_text);
75 }
76
77 let exa_response: ExaSearchResponse = response.json().await?;
78
79 let mut results = SearchResults::new(query.to_string(), "exa".to_string());
80 results.set_search_time(search_time_ms);
81
82 for exa_result in exa_response.results {
83 let snippet = exa_result.text.unwrap_or_else(|| exa_result.title.clone());
85
86 let search_result = SearchResult {
87 title: exa_result.title,
88 url: exa_result.url,
89 snippet,
90 published_date: exa_result.published_date,
91 author: exa_result.author,
92 score: exa_result.score.map(|s| s as f32),
93 };
94
95 results.add_result(search_result);
96 }
97
98 Ok(results)
99}
100
101#[cfg(test)]
102mod tests {
103 use super::*;
104
105 #[test]
106 fn test_exa_response_parsing() {
107 let json_response = r#"{
108 "results": [
109 {
110 "title": "Understanding AI Safety",
111 "url": "https://example.com/ai-safety",
112 "text": "AI safety is a critical field of research...",
113 "publishedDate": "2024-01-15",
114 "author": "Jane Doe",
115 "score": 0.95
116 }
117 ]
118 }"#;
119
120 let response: ExaSearchResponse = serde_json::from_str(json_response).unwrap();
121 assert_eq!(response.results.len(), 1);
122 assert_eq!(response.results[0].title, "Understanding AI Safety");
123 assert_eq!(response.results[0].score, Some(0.95));
124 }
125}