Skip to main content

fastpaper/sources/
semantic.rs

1use std::sync::{Mutex, OnceLock};
2use std::time::{Duration, Instant};
3
4use super::Paper;
5
6const FIELDS: &str = "title,abstract,year,citationCount,authors,url,publicationDate,externalIds,fieldsOfStudy,openAccessPdf,venue";
7
8const USER_AGENT: &str = concat!(
9    "fastpaper-cli/",
10    env!("CARGO_PKG_VERSION"),
11    " (+https://github.com/zhangyee/fastpaper-cli)"
12);
13
14// 节流间隔(带 key)。
15const MIN_INTERVAL_AUTH: Duration = Duration::from_millis(1000);
16// 节流间隔(匿名)。
17const MIN_INTERVAL_ANON: Duration = Duration::from_millis(100);
18const MAX_RETRIES: u32 = 5;
19
20#[derive(Clone, Copy)]
21struct BackoffConfig {
22    base: Duration,
23    max: Duration,
24    max_retries: u32,
25}
26
27impl BackoffConfig {
28    const DEFAULT_AUTH: Self = Self {
29        base: Duration::from_secs(1),
30        max: Duration::from_secs(30),
31        max_retries: MAX_RETRIES,
32    };
33    const DEFAULT_ANON: Self = Self {
34        base: Duration::from_secs(2),
35        max: Duration::from_secs(30),
36        max_retries: MAX_RETRIES,
37    };
38}
39
40enum FetchOutcome {
41    Ok(String),
42    RateLimited,
43    Err(String),
44}
45
46static LAST_CALL: OnceLock<Mutex<Option<Instant>>> = OnceLock::new();
47
48// 解析 HTTP Retry-After 头。
49fn parse_retry_after(headers: &ureq::http::HeaderMap, max: Duration) -> Option<Duration> {
50    headers
51        .get("retry-after")
52        .and_then(|v| v.to_str().ok())
53        .and_then(|s| s.trim().parse::<u64>().ok())
54        .map(Duration::from_secs)
55        .map(|d| d.min(max))
56}
57
58// 指数退避(带 0-10% jitter)。
59fn backoff_delay(attempt: u32, cfg: &BackoffConfig) -> Duration {
60    let shift = attempt.min(10);
61    let exp_ms = (cfg.base.as_millis() as u64).saturating_mul(1u64 << shift);
62    let exp = Duration::from_millis(exp_ms).min(cfg.max);
63    // 轻量 jitter:用 Instant nanos 做伪随机 0..=10%
64    let nanos = Instant::now().elapsed().subsec_nanos() as u64;
65    let jitter_pct = nanos % 11; // 0..=10
66    let jitter_ms = (exp.as_millis() as u64).saturating_mul(jitter_pct) / 100;
67    exp + Duration::from_millis(jitter_ms)
68}
69
70// 进程级节流。第二次连续调用会 sleep 到至少 min_interval。
71fn throttle_with(state: &Mutex<Option<Instant>>, min_interval: Duration) {
72    let mut guard = state.lock().unwrap();
73    if let Some(last) = *guard {
74        let elapsed = last.elapsed();
75        if elapsed < min_interval {
76            std::thread::sleep(min_interval - elapsed);
77        }
78    }
79    *guard = Some(Instant::now());
80}
81
82fn throttle(has_key: bool) {
83    let min = if has_key {
84        MIN_INTERVAL_AUTH
85    } else {
86        MIN_INTERVAL_ANON
87    };
88    let state = LAST_CALL.get_or_init(|| Mutex::new(None));
89    throttle_with(state, min);
90}
91
92fn http_get_with_retry_cfg(
93    url: &str,
94    mut api_key: Option<String>,
95    cfg: &BackoffConfig,
96) -> FetchOutcome {
97    let mut tried_without_key = api_key.is_none();
98    let agent: ureq::Agent = ureq::Agent::config_builder()
99        .http_status_as_error(false)
100        .build()
101        .into();
102
103    let mut attempt: u32 = 0;
104    while attempt <= cfg.max_retries {
105        throttle(api_key.is_some());
106
107        let mut req = agent.get(url).header("User-Agent", USER_AGENT);
108        if let Some(ref k) = api_key {
109            req = req.header("x-api-key", k);
110        }
111        let resp = match req.call() {
112            Ok(r) => r,
113            Err(e) => return FetchOutcome::Err(format!("HTTP error: {e}")),
114        };
115        let status = resp.status().as_u16();
116        match status {
117            200 => {
118                return match resp.into_body().read_to_string() {
119                    Ok(body) => FetchOutcome::Ok(body),
120                    Err(e) => FetchOutcome::Err(format!("read body: {e}")),
121                };
122            }
123            403 if api_key.is_some() && !tried_without_key => {
124                eprintln!("[semantic] API key rejected (403), retrying without key");
125                api_key = None;
126                tried_without_key = true;
127                continue; // 不计 attempt
128            }
129            404 => return FetchOutcome::Err("HTTP 404".to_string()),
130            429 => {
131                if attempt == cfg.max_retries {
132                    return FetchOutcome::RateLimited;
133                }
134                let wait = parse_retry_after(resp.headers(), cfg.max)
135                    .unwrap_or_else(|| backoff_delay(attempt, cfg));
136                std::thread::sleep(wait);
137            }
138            500..=599 => return FetchOutcome::Err(format!("Server error: {status}")),
139            _ => return FetchOutcome::Err(format!("HTTP {status}")),
140        }
141        attempt += 1;
142    }
143    FetchOutcome::RateLimited
144}
145
146/// Relevance search pages with offset only this far.
147const MAX_OFFSET: u32 = 1000;
148
149/// Build the search URL.
150///
151/// Two endpoints: `/paper/search` ranks by relevance but cannot sort, and
152/// `/paper/search/bulk` sorts but pages with an opaque token instead of an
153/// offset. Asking for an explicit order switches to bulk.
154fn build_search_url(base_url: &str, q: &super::SearchQuery) -> Result<String, String> {
155    if let Some(ref author) = q.author {
156        return Err(format!(
157            "semantic has no author filter on paper search (asked for --author {}).\n\
158             Try arxiv, crossref or openalex, which do.",
159            author
160        ));
161    }
162
163    let sort = match q.sort {
164        // Relevance is what /paper/search already does.
165        None | Some(super::SortField::Relevance) => None,
166        Some(super::SortField::Date) => Some("publicationDate"),
167        Some(super::SortField::Citations) => Some("citationCount"),
168    };
169
170    let mut url = if sort.is_some() {
171        if q.offset > 0 {
172            return Err(
173                "semantic pages sorted results with a token, not an offset, so --offset cannot be \
174                 combined with --sort date or --sort citations"
175                    .to_string(),
176            );
177        }
178        format!(
179            "{}/graph/v1/paper/search/bulk?query={}&fields={}",
180            base_url,
181            super::encode_query(&q.query),
182            FIELDS
183        )
184    } else {
185        if q.offset > MAX_OFFSET {
186            return Err(format!(
187                "semantic supports an offset up to {} (asked for {})",
188                MAX_OFFSET, q.offset
189            ));
190        }
191        format!(
192            "{}/graph/v1/paper/search?query={}&limit={}&offset={}&fields={}",
193            base_url,
194            super::encode_query(&q.query),
195            q.limit,
196            q.offset,
197            FIELDS
198        )
199    };
200
201    if let Some(by) = sort {
202        let order = match q.order {
203            super::SortOrder::Asc => "asc",
204            super::SortOrder::Desc => "desc",
205        };
206        url.push_str(&format!("&sort={}:{}", by, order));
207    }
208
209    match q.year {
210        Some(year) => url.push_str(&format!("&year={}", year)),
211        None => {
212            if q.after.is_some() || q.before.is_some() {
213                let from = match q.after {
214                    Some(ref d) => super::validate_ymd(d)?,
215                    None => "",
216                };
217                let to = match q.before {
218                    Some(ref d) => super::validate_ymd(d)?,
219                    None => "",
220                };
221                url.push_str(&format!("&publicationDateOrYear={}:{}", from, to));
222            }
223        }
224    }
225
226    if let Some(ref field) = q.field {
227        url.push_str(&format!("&fieldsOfStudy={}", super::encode_query(field)));
228    }
229    // openAccessPdf is a presence flag; giving it a value is not the documented
230    // form.
231    if q.open_access {
232        url.push_str("&openAccessPdf");
233    }
234
235    Ok(url)
236}
237
238/// Search Semantic Scholar API and return parsed papers.
239pub fn search(base_url: &str, q: &super::SearchQuery) -> Result<Vec<Paper>, String> {
240    let url = build_search_url(base_url, q)?;
241    // The bulk endpoint takes no limit; it returns a full page plus a token.
242    let truncate_to = if q.sort.is_some() && q.sort != Some(super::SortField::Relevance) {
243        Some(q.limit as usize)
244    } else {
245        None
246    };
247    let api_key = std::env::var("SEMANTIC_SCHOLAR_API_KEY").ok();
248    let cfg = if api_key.is_some() {
249        BackoffConfig::DEFAULT_AUTH
250    } else {
251        BackoffConfig::DEFAULT_ANON
252    };
253    match http_get_with_retry_cfg(&url, api_key, &cfg) {
254        FetchOutcome::Ok(body) => {
255            let mut papers = parse_search_response(&body)?;
256            if let Some(n) = truncate_to {
257                papers.truncate(n);
258            }
259            Ok(papers)
260        }
261        FetchOutcome::RateLimited => Err(format!("rate limited after {} retries", cfg.max_retries)),
262        FetchOutcome::Err(e) => Err(e),
263    }
264}
265
266#[cfg(test)]
267fn search_with_cfg_for_test(
268    base_url: &str,
269    query: &str,
270    max_results: u32,
271    cfg: &BackoffConfig,
272) -> Result<Vec<Paper>, String> {
273    let encoded = super::encode_query(query);
274    let url = format!(
275        "{}/graph/v1/paper/search?query={}&limit={}&fields={}",
276        base_url, encoded, max_results, FIELDS
277    );
278    let api_key = std::env::var("SEMANTIC_SCHOLAR_API_KEY").ok();
279    match http_get_with_retry_cfg(&url, api_key, cfg) {
280        FetchOutcome::Ok(body) => parse_search_response(&body),
281        FetchOutcome::RateLimited => Err(format!("rate limited after {} retries", cfg.max_retries)),
282        FetchOutcome::Err(e) => Err(e),
283    }
284}
285
286/// Fetch a single paper by S2 paper ID.
287pub fn get_by_id(base_url: &str, s2_id: &str) -> Result<Option<Paper>, String> {
288    let url = format!("{}/graph/v1/paper/{}?fields={}", base_url, s2_id, FIELDS);
289    let api_key = std::env::var("SEMANTIC_SCHOLAR_API_KEY").ok();
290    let cfg = if api_key.is_some() {
291        BackoffConfig::DEFAULT_AUTH
292    } else {
293        BackoffConfig::DEFAULT_ANON
294    };
295    match http_get_with_retry_cfg(&url, api_key, &cfg) {
296        FetchOutcome::Ok(body) => {
297            let wrapped = format!(r#"{{"data":[{}]}}"#, body);
298            Ok(parse_search_response(&wrapped)?.into_iter().next())
299        }
300        FetchOutcome::Err(e) if e.contains("404") => Ok(None),
301        FetchOutcome::Err(e) => Err(e),
302        FetchOutcome::RateLimited => Err(format!("rate limited after {} retries", cfg.max_retries)),
303    }
304}
305
306#[cfg(test)]
307fn get_by_id_with_cfg_for_test(
308    base_url: &str,
309    s2_id: &str,
310    cfg: &BackoffConfig,
311) -> Result<Option<Paper>, String> {
312    let url = format!("{}/graph/v1/paper/{}?fields={}", base_url, s2_id, FIELDS);
313    let api_key = std::env::var("SEMANTIC_SCHOLAR_API_KEY").ok();
314    match http_get_with_retry_cfg(&url, api_key, cfg) {
315        FetchOutcome::Ok(body) => {
316            let wrapped = format!(r#"{{"data":[{}]}}"#, body);
317            Ok(parse_search_response(&wrapped)?.into_iter().next())
318        }
319        FetchOutcome::Err(e) if e.contains("404") => Ok(None),
320        FetchOutcome::Err(e) => Err(e),
321        FetchOutcome::RateLimited => Err(format!("rate limited after {} retries", cfg.max_retries)),
322    }
323}
324
325/// Walk a citation edge. Both directions are one request.
326pub fn cite(
327    base_url: &str,
328    id: &str,
329    direction: super::Direction,
330    limit: u32,
331) -> Result<Vec<Paper>, String> {
332    let endpoint = match direction {
333        super::Direction::Incoming => "citations",
334        super::Direction::Outgoing => "references",
335    };
336    let url = format!(
337        "{}/graph/v1/paper/{}/{}?limit={}&fields={}",
338        base_url, id, endpoint, limit, FIELDS
339    );
340    let key = std::env::var("SEMANTIC_SCHOLAR_API_KEY").ok();
341    let cfg = if key.is_some() {
342        BackoffConfig::DEFAULT_AUTH
343    } else {
344        BackoffConfig::DEFAULT_ANON
345    };
346    match http_get_with_retry_cfg(&url, key, &cfg) {
347        FetchOutcome::Ok(body) => parse_edge_response(&body, direction),
348        FetchOutcome::RateLimited => Err(
349            "semantic rate limited. Set SEMANTIC_SCHOLAR_API_KEY, or use: fastpaper cite openalex <id>"
350                .to_string(),
351        ),
352        FetchOutcome::Err(e) => Err(e),
353    }
354}
355
356/// Parse a `/citations` or `/references` response.
357///
358/// Both wrap each paper one level deeper than `/paper/search` does — under
359/// `citingPaper` coming in, `citedPaper` going out — but the paper object
360/// itself is identical, so the mapping is shared.
361pub fn parse_edge_response(json: &str, direction: super::Direction) -> Result<Vec<Paper>, String> {
362    let root: serde_json::Value =
363        serde_json::from_str(json).map_err(|e| format!("JSON parse error: {}", e))?;
364    let data = root["data"].as_array().ok_or("missing 'data' array")?;
365
366    let key = match direction {
367        super::Direction::Incoming => "citingPaper",
368        super::Direction::Outgoing => "citedPaper",
369    };
370
371    Ok(data
372        .iter()
373        .map(|edge| &edge[key])
374        .filter(|p| !p.is_null())
375        .map(paper_from)
376        .filter(|p| !p.title.is_empty())
377        .collect())
378}
379
380/// Parse Semantic Scholar JSON search response into a list of Papers.
381pub fn parse_search_response(json: &str) -> Result<Vec<Paper>, String> {
382    let root: serde_json::Value =
383        serde_json::from_str(json).map_err(|e| format!("JSON parse error: {}", e))?;
384
385    let data = root["data"].as_array().ok_or("missing 'data' array")?;
386
387    Ok(data.iter().map(paper_from).collect())
388}
389
390/// Map one Semantic Scholar paper object onto `Paper`. Shared by search,
391/// single-paper lookup and both citation-edge endpoints.
392fn paper_from(item: &serde_json::Value) -> Paper {
393    let authors: Vec<String> = item["authors"]
394        .as_array()
395        .map(|arr| {
396            arr.iter()
397                .filter_map(|a| a["name"].as_str().map(|s| s.to_string()))
398                .collect()
399        })
400        .unwrap_or_default();
401
402    let doi = item["externalIds"]["DOI"].as_str().map(|s| s.to_string());
403
404    let pdf_url = item["openAccessPdf"]["url"]
405        .as_str()
406        .filter(|s| !s.is_empty())
407        .map(|s| {
408            if s.contains("arxiv.org/abs/") {
409                s.replace("/abs/", "/pdf/")
410            } else {
411                s.to_string()
412            }
413        });
414
415    let fields: Vec<String> = item["fieldsOfStudy"]
416        .as_array()
417        .map(|arr| {
418            arr.iter()
419                .filter_map(|v| v.as_str().map(|s| s.to_string()))
420                .collect()
421        })
422        .unwrap_or_default();
423
424    let citations = item["citationCount"].as_u64().map(|n| n as u32);
425
426    Paper {
427        id: item["paperId"].as_str().unwrap_or("").to_string(),
428        title: item["title"].as_str().unwrap_or("").to_string(),
429        authors,
430        abstract_text: item["abstract"].as_str().map(|s| s.to_string()),
431        year: item["year"].as_u64().map(|y| y as u16),
432        doi,
433        url: item["url"].as_str().map(|s| s.to_string()),
434        pdf_url,
435        venue: item["venue"].as_str().map(|s| s.to_string()),
436        citations,
437        fields,
438        open_access: Some(item["openAccessPdf"].is_object() && !item["openAccessPdf"].is_null()),
439        source: "semantic".to_string(),
440    }
441}
442
443#[cfg(test)]
444mod tests {
445    use super::*;
446    use serial_test::serial;
447
448    const FIXTURE: &str = include_str!("../../tests/fixtures/semantic_search.json");
449
450    #[test]
451    fn parse_returns_ok() {
452        let result = parse_search_response(FIXTURE);
453        assert!(result.is_ok());
454    }
455
456    #[test]
457    fn parse_papers_not_empty() {
458        let papers = parse_search_response(FIXTURE).unwrap();
459        assert!(!papers.is_empty());
460    }
461
462    #[test]
463    fn parse_titles_not_empty() {
464        let papers = parse_search_response(FIXTURE).unwrap();
465        for p in &papers {
466            assert!(!p.title.is_empty(), "paper {} has empty title", p.id);
467        }
468    }
469
470    #[test]
471    fn parse_source_is_semantic() {
472        let papers = parse_search_response(FIXTURE).unwrap();
473        for p in &papers {
474            assert_eq!(p.source, "semantic");
475        }
476    }
477
478    #[test]
479    fn parse_citations_present() {
480        let papers = parse_search_response(FIXTURE).unwrap();
481        for p in &papers {
482            assert!(p.citations.is_some(), "paper {} missing citations", p.id);
483            assert!(p.citations.unwrap() > 0);
484        }
485    }
486
487    #[test]
488    fn parse_pdf_url_from_open_access() {
489        // Our fixture has openAccessPdf as null for all papers
490        let papers = parse_search_response(FIXTURE).unwrap();
491        for p in &papers {
492            assert!(p.pdf_url.is_none(), "paper {} should have no pdf_url", p.id);
493        }
494    }
495
496    #[test]
497    fn parse_doi_from_external_ids() {
498        let papers = parse_search_response(FIXTURE).unwrap();
499        // First paper has DOI in externalIds
500        let first = &papers[0];
501        assert_eq!(first.doi.as_deref(), Some("10.1016/J.NEUCOM.2021.03.091"));
502    }
503
504    #[test]
505    fn parse_empty_data_returns_empty_list() {
506        let papers = parse_search_response(r#"{"data": []}"#).unwrap();
507        assert!(papers.is_empty());
508    }
509
510    #[test]
511    fn search_returns_papers() {
512        let mut server = mockito::Server::new();
513        let mock = server
514            .mock("GET", mockito::Matcher::Any)
515            .with_status(200)
516            .with_body(FIXTURE)
517            .create();
518        let papers = search(
519            &server.url(),
520            &crate::sources::SearchQuery::simple("test", 3),
521        )
522        .unwrap();
523        assert!(!papers.is_empty());
524        mock.assert();
525    }
526
527    #[test]
528    fn search_request_path_contains_paper_search() {
529        let mut server = mockito::Server::new();
530        let mock = server
531            .mock("GET", mockito::Matcher::Regex("paper/search".to_string()))
532            .with_status(200)
533            .with_body(FIXTURE)
534            .create();
535        let _ = search(
536            &server.url(),
537            &crate::sources::SearchQuery::simple("test", 3),
538        );
539        mock.assert();
540    }
541
542    #[test]
543    fn search_request_contains_query_param() {
544        let mut server = mockito::Server::new();
545        let mock = server
546            .mock("GET", mockito::Matcher::Regex("query=test".to_string()))
547            .with_status(200)
548            .with_body(FIXTURE)
549            .create();
550        let _ = search(
551            &server.url(),
552            &crate::sources::SearchQuery::simple("test", 3),
553        );
554        mock.assert();
555    }
556
557    #[test]
558    fn search_request_contains_limit() {
559        let mut server = mockito::Server::new();
560        let mock = server
561            .mock("GET", mockito::Matcher::Regex("limit=3".to_string()))
562            .with_status(200)
563            .with_body(FIXTURE)
564            .create();
565        let _ = search(
566            &server.url(),
567            &crate::sources::SearchQuery::simple("test", 3),
568        );
569        mock.assert();
570    }
571
572    #[test]
573    #[serial]
574    fn search_sends_api_key_header_when_set() {
575        let mut server = mockito::Server::new();
576        let mock = server
577            .mock("GET", mockito::Matcher::Any)
578            .match_header("x-api-key", "test-key-123")
579            .with_status(200)
580            .with_body(FIXTURE)
581            .create();
582        unsafe { std::env::set_var("SEMANTIC_SCHOLAR_API_KEY", "test-key-123") };
583        let _ = search(
584            &server.url(),
585            &crate::sources::SearchQuery::simple("test", 3),
586        );
587        unsafe { std::env::remove_var("SEMANTIC_SCHOLAR_API_KEY") };
588        mock.assert();
589    }
590
591    #[test]
592    #[serial]
593    fn search_works_without_api_key() {
594        unsafe { std::env::remove_var("SEMANTIC_SCHOLAR_API_KEY") };
595        let mut server = mockito::Server::new();
596        let mock = server
597            .mock("GET", mockito::Matcher::Any)
598            .with_status(200)
599            .with_body(FIXTURE)
600            .create();
601        let result = search(
602            &server.url(),
603            &crate::sources::SearchQuery::simple("test", 3),
604        );
605        assert!(result.is_ok());
606        mock.assert();
607    }
608
609    #[test]
610    fn parse_retry_after_seconds() {
611        let cap = Duration::from_secs(30);
612
613        // helper 构造一个只含 retry-after 的 HeaderMap
614        fn headers_with_retry_after(value: &str) -> ureq::http::HeaderMap {
615            let mut h = ureq::http::HeaderMap::new();
616            h.insert(
617                ureq::http::header::HeaderName::from_static("retry-after"),
618                ureq::http::HeaderValue::from_str(value).unwrap(),
619            );
620            h
621        }
622
623        assert_eq!(
624            parse_retry_after(&headers_with_retry_after("5"), cap),
625            Some(Duration::from_secs(5))
626        );
627        assert_eq!(
628            parse_retry_after(&headers_with_retry_after("abc"), cap),
629            None
630        );
631        assert_eq!(parse_retry_after(&ureq::http::HeaderMap::new(), cap), None);
632        assert_eq!(
633            parse_retry_after(&headers_with_retry_after("9999"), cap),
634            Some(cap)
635        );
636    }
637
638    #[test]
639    fn backoff_delay_grows_exponentially() {
640        let cfg = BackoffConfig {
641            base: Duration::from_millis(100),
642            max: Duration::from_secs(5),
643            max_retries: 5,
644        };
645        let d0 = backoff_delay(0, &cfg);
646        let d2 = backoff_delay(2, &cfg);
647        assert!(d0 < d2, "delay should grow with attempt");
648        assert!(
649            d2 <= cfg.max + cfg.max / 10,
650            "delay should respect max (with up-to-10% jitter)"
651        );
652    }
653
654    #[test]
655    fn backoff_delay_does_not_overflow() {
656        let cfg = BackoffConfig {
657            base: Duration::from_secs(1),
658            max: Duration::from_secs(30),
659            max_retries: 5,
660        };
661        // attempt = 20 不应导致左移溢出
662        let d = backoff_delay(20, &cfg);
663        assert!(d <= cfg.max + cfg.max / 10);
664    }
665
666    #[test]
667    fn throttle_with_sleeps_when_called_back_to_back() {
668        let state: Mutex<Option<Instant>> = Mutex::new(None);
669        let interval = Duration::from_millis(80);
670
671        let t0 = Instant::now();
672        throttle_with(&state, interval);
673        throttle_with(&state, interval);
674        let elapsed = t0.elapsed();
675
676        assert!(
677            elapsed >= interval,
678            "second call should wait at least one interval (got {:?})",
679            elapsed
680        );
681    }
682
683    #[test]
684    fn throttle_with_does_not_sleep_after_long_gap() {
685        let past = Instant::now() - Duration::from_secs(3600);
686        let state: Mutex<Option<Instant>> = Mutex::new(Some(past));
687        let interval = Duration::from_millis(100);
688
689        let t0 = Instant::now();
690        throttle_with(&state, interval);
691        let elapsed = t0.elapsed();
692
693        assert!(
694            elapsed < Duration::from_millis(20),
695            "should return immediately after long gap (got {:?})",
696            elapsed
697        );
698    }
699
700    #[test]
701    #[serial]
702    fn request_sends_user_agent_header() {
703        unsafe { std::env::remove_var("SEMANTIC_SCHOLAR_API_KEY") };
704        let mut server = mockito::Server::new();
705        let mock = server
706            .mock("GET", mockito::Matcher::Any)
707            .match_header(
708                "user-agent",
709                mockito::Matcher::Regex(r"^fastpaper-cli/\d+\.\d+\.\d+ \(\+https://".to_string()),
710            )
711            .with_status(200)
712            .with_body(FIXTURE)
713            .create();
714        let _ = search(
715            &server.url(),
716            &crate::sources::SearchQuery::simple("test", 3),
717        );
718        mock.assert();
719    }
720
721    #[test]
722    #[serial]
723    fn search_returns_err_on_rate_limit_exhausted() {
724        unsafe { std::env::remove_var("SEMANTIC_SCHOLAR_API_KEY") };
725        let mut server = mockito::Server::new();
726        let _m = server
727            .mock("GET", mockito::Matcher::Any)
728            .with_status(429)
729            .expect_at_least(4)
730            .create();
731
732        let cfg = BackoffConfig {
733            base: Duration::from_millis(10),
734            max: Duration::from_millis(50),
735            max_retries: 3,
736        };
737        let result = search_with_cfg_for_test(&server.url(), "test", 3, &cfg);
738
739        assert!(
740            matches!(result, Err(ref e) if e.contains("rate limited")),
741            "expected Err on rate-limit exhausted (consistent with other sources), got {:?}",
742            result
743        );
744    }
745
746    #[test]
747    #[serial]
748    fn rate_limit_then_success_respects_retry_after() {
749        unsafe { std::env::remove_var("SEMANTIC_SCHOLAR_API_KEY") };
750        let mut server = mockito::Server::new();
751
752        // 第 1 次:429 + Retry-After: 1
753        let m1 = server
754            .mock("GET", mockito::Matcher::Any)
755            .with_status(429)
756            .with_header("retry-after", "1")
757            .expect(1)
758            .create();
759        // 第 2 次:200
760        let m2 = server
761            .mock("GET", mockito::Matcher::Any)
762            .with_status(200)
763            .with_body(FIXTURE)
764            .expect(1)
765            .create();
766
767        let t0 = Instant::now();
768        let result = search(
769            &server.url(),
770            &crate::sources::SearchQuery::simple("test", 3),
771        );
772        let elapsed = t0.elapsed();
773
774        assert!(result.is_ok(), "expected Ok, got {:?}", result);
775        assert!(!result.unwrap().is_empty());
776        assert!(
777            elapsed >= Duration::from_secs(1) && elapsed < Duration::from_secs(3),
778            "expected ~1s wait (got {:?})",
779            elapsed
780        );
781        m1.assert();
782        m2.assert();
783    }
784
785    #[test]
786    #[serial]
787    fn get_by_id_returns_err_on_rate_limit_exhausted() {
788        unsafe { std::env::remove_var("SEMANTIC_SCHOLAR_API_KEY") };
789        let mut server = mockito::Server::new();
790        let _m = server
791            .mock("GET", mockito::Matcher::Any)
792            .with_status(429)
793            .expect_at_least(4)
794            .create();
795
796        let cfg = BackoffConfig {
797            base: Duration::from_millis(10),
798            max: Duration::from_millis(50),
799            max_retries: 3,
800        };
801        let result = get_by_id_with_cfg_for_test(&server.url(), "abc123", &cfg);
802
803        assert!(
804            matches!(result, Err(ref e) if e.contains("rate limited")),
805            "expected Err on rate-limit exhausted, got {:?}",
806            result
807        );
808    }
809
810    #[test]
811    #[serial]
812    fn forbidden_with_key_strips_key_and_retries() {
813        unsafe { std::env::set_var("SEMANTIC_SCHOLAR_API_KEY", "bad-key") };
814
815        let mut server = mockito::Server::new();
816        // 带 key → 403
817        let m_403 = server
818            .mock("GET", mockito::Matcher::Any)
819            .match_header("x-api-key", "bad-key")
820            .with_status(403)
821            .expect(1)
822            .create();
823        // 不带 key → 200
824        let m_200 = server
825            .mock("GET", mockito::Matcher::Any)
826            .match_header("x-api-key", mockito::Matcher::Missing)
827            .with_status(200)
828            .with_body(FIXTURE)
829            .expect(1)
830            .create();
831
832        let result = search(
833            &server.url(),
834            &crate::sources::SearchQuery::simple("test", 3),
835        );
836        unsafe { std::env::remove_var("SEMANTIC_SCHOLAR_API_KEY") };
837
838        assert!(
839            result.is_ok(),
840            "expected Ok after key-strip retry, got {:?}",
841            result
842        );
843        assert!(!result.unwrap().is_empty());
844        m_403.assert();
845        m_200.assert();
846    }
847}
848
849#[cfg(test)]
850mod query_tests {
851    use super::*;
852    use crate::sources::{SearchQuery, SortField, SortOrder};
853
854    fn url(q: &SearchQuery) -> String {
855        build_search_url("https://api.semanticscholar.org", q).unwrap()
856    }
857
858    #[test]
859    fn plain_query_uses_the_relevance_endpoint() {
860        let u = url(&SearchQuery::simple("attention", 10));
861        assert!(u.contains("/graph/v1/paper/search?"), "got: {}", u);
862        assert!(u.contains("query=attention"), "got: {}", u);
863        assert!(u.contains("limit=10"), "got: {}", u);
864    }
865
866    #[test]
867    fn offset_is_passed_through() {
868        let mut q = SearchQuery::simple("attention", 10);
869        q.offset = 40;
870        assert!(url(&q).contains("offset=40"));
871    }
872
873    #[test]
874    fn offset_beyond_the_api_limit_is_rejected() {
875        let mut q = SearchQuery::simple("attention", 10);
876        q.offset = 1001;
877        let err = build_search_url("https://api.semanticscholar.org", &q).unwrap_err();
878        assert!(err.contains("1000"), "got: {}", err);
879    }
880
881    #[test]
882    fn year_is_passed_through() {
883        let mut q = SearchQuery::simple("attention", 10);
884        q.year = Some(2017);
885        assert!(url(&q).contains("year=2017"), "got: {}", url(&q));
886    }
887
888    #[test]
889    fn dates_become_a_publication_date_range() {
890        let mut q = SearchQuery::simple("attention", 10);
891        q.after = Some("2024-01-01".into());
892        q.before = Some("2024-03-31".into());
893        assert!(
894            url(&q).contains("publicationDateOrYear=2024-01-01:2024-03-31"),
895            "got: {}",
896            url(&q)
897        );
898    }
899
900    #[test]
901    fn after_alone_leaves_the_upper_bound_open() {
902        let mut q = SearchQuery::simple("attention", 10);
903        q.after = Some("2024-01-01".into());
904        assert!(
905            url(&q).contains("publicationDateOrYear=2024-01-01:"),
906            "got: {}",
907            url(&q)
908        );
909    }
910
911    #[test]
912    fn field_becomes_fields_of_study() {
913        let mut q = SearchQuery::simple("attention", 10);
914        q.field = Some("Computer Science".into());
915        assert!(
916            url(&q).contains("fieldsOfStudy=Computer"),
917            "got: {}",
918            url(&q)
919        );
920    }
921
922    // openAccessPdf is a valueless presence flag.
923    #[test]
924    fn open_access_adds_a_bare_flag() {
925        let mut q = SearchQuery::simple("attention", 10);
926        q.open_access = true;
927        let u = url(&q);
928        assert!(u.contains("openAccessPdf"), "got: {}", u);
929        assert!(
930            !u.contains("openAccessPdf=true"),
931            "flag takes no value: {}",
932            u
933        );
934    }
935
936    // The bulk endpoint has no limit parameter -- it answers with up to a
937    // thousand records and a continuation token -- so -n has to be applied
938    // after the fact or `--sort citations -n 3` prints the whole page.
939    #[test]
940    fn bulk_results_are_truncated_to_the_limit() {
941        let body = format!(
942            r#"{{"total":50,"data":[{}]}}"#,
943            (0..50)
944                .map(|i| format!(r#"{{"paperId":"p{}","title":"T{}"}}"#, i, i))
945                .collect::<Vec<_>>()
946                .join(",")
947        );
948        let mut server = mockito::Server::new();
949        server
950            .mock("GET", mockito::Matcher::Any)
951            .with_status(200)
952            .with_body(body)
953            .create();
954        let mut q = SearchQuery::simple("x", 3);
955        q.sort = Some(SortField::Citations);
956        let papers = search(&server.url(), &q).unwrap();
957        assert_eq!(papers.len(), 3, "should honour -n on the bulk endpoint");
958    }
959
960    // Only the bulk endpoint sorts; relevance ranking is the default elsewhere.
961    #[test]
962    fn sorting_switches_to_the_bulk_endpoint() {
963        let mut q = SearchQuery::simple("attention", 10);
964        q.sort = Some(SortField::Citations);
965        let u = url(&q);
966        assert!(u.contains("/paper/search/bulk?"), "got: {}", u);
967        assert!(u.contains("sort=citationCount:desc"), "got: {}", u);
968    }
969
970    #[test]
971    fn sort_by_date_uses_publication_date() {
972        let mut q = SearchQuery::simple("attention", 10);
973        q.sort = Some(SortField::Date);
974        q.order = SortOrder::Asc;
975        assert!(
976            url(&q).contains("sort=publicationDate:asc"),
977            "got: {}",
978            url(&q)
979        );
980    }
981
982    #[test]
983    fn sort_by_relevance_stays_on_the_relevance_endpoint() {
984        let mut q = SearchQuery::simple("attention", 10);
985        q.sort = Some(SortField::Relevance);
986        let u = url(&q);
987        assert!(u.contains("/paper/search?"), "got: {}", u);
988        assert!(
989            !u.contains("sort="),
990            "relevance is the default ranking: {}",
991            u
992        );
993    }
994
995    // The bulk endpoint has a token cursor instead of offset.
996    #[test]
997    fn offset_with_sorting_is_rejected() {
998        let mut q = SearchQuery::simple("attention", 10);
999        q.sort = Some(SortField::Citations);
1000        q.offset = 10;
1001        let err = build_search_url("https://api.semanticscholar.org", &q).unwrap_err();
1002        assert!(err.contains("--offset"), "got: {}", err);
1003    }
1004
1005    // Semantic Scholar's paper search has no author parameter.
1006    #[test]
1007    fn author_is_rejected_with_an_alternative() {
1008        let mut q = SearchQuery::simple("attention", 10);
1009        q.author = Some("Vaswani".into());
1010        let err = build_search_url("https://api.semanticscholar.org", &q).unwrap_err();
1011        assert!(err.contains("--author"), "got: {}", err);
1012    }
1013}
1014
1015#[cfg(test)]
1016mod edge_tests {
1017    use super::*;
1018    use crate::sources::Direction;
1019
1020    const REFS: &str = include_str!("../../tests/fixtures/semantic_references.json");
1021    const CITES: &str = include_str!("../../tests/fixtures/semantic_citations.json");
1022
1023    // The edge endpoints wrap each paper one level deeper than /paper/search:
1024    // `data[].citedPaper` going out, `data[].citingPaper` coming in.
1025    #[test]
1026    fn outgoing_edges_unwrap_cited_papers() {
1027        let papers = parse_edge_response(REFS, Direction::Outgoing).unwrap();
1028        assert_eq!(papers.len(), 3);
1029        assert!(papers.iter().all(|p| !p.title.is_empty()));
1030        assert!(papers.iter().all(|p| p.source == "semantic"));
1031    }
1032
1033    #[test]
1034    fn incoming_edges_unwrap_citing_papers() {
1035        let papers = parse_edge_response(CITES, Direction::Incoming).unwrap();
1036        assert_eq!(papers.len(), 3);
1037        assert!(papers.iter().all(|p| !p.title.is_empty()));
1038    }
1039
1040    // Reading the wrong side yields nothing rather than silently mixing
1041    // directions, which would invert every edge in a citation graph.
1042    #[test]
1043    fn reading_the_wrong_side_yields_no_papers() {
1044        assert!(
1045            parse_edge_response(REFS, Direction::Incoming)
1046                .unwrap()
1047                .is_empty()
1048        );
1049    }
1050
1051    #[test]
1052    fn edge_papers_carry_the_same_fields_as_search_results() {
1053        let p = &parse_edge_response(CITES, Direction::Incoming).unwrap()[0];
1054        assert!(p.year.is_some(), "year missing");
1055        assert!(p.citations.is_some(), "citationCount missing");
1056        assert!(!p.authors.is_empty(), "authors missing");
1057    }
1058}