Skip to main content

lc/search/
exa.rs

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    // The provider_config.url should be the complete search endpoint URL
51    // For Exa, it should be https://api.exa.ai/search
52    let base_url = provider_config.url.trim_end_matches('/');
53    let url = if base_url.ends_with("/search") {
54        // URL already includes the endpoint path
55        base_url.to_string()
56    } else {
57        // URL is just the base, append the endpoint
58        format!("{}/search", base_url)
59    };
60    let mut request = client.post(&url).json(&request_body);
61
62    // Add headers
63    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        // Use text content as snippet, or title if text is not available
84        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}