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
14const MIN_INTERVAL_AUTH: Duration = Duration::from_millis(1000);
16const 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
48fn 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
58fn 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 let nanos = Instant::now().elapsed().subsec_nanos() as u64;
65 let jitter_pct = nanos % 11; let jitter_ms = (exp.as_millis() as u64).saturating_mul(jitter_pct) / 100;
67 exp + Duration::from_millis(jitter_ms)
68}
69
70fn 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; }
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
146const MAX_OFFSET: u32 = 1000;
148
149fn 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 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 if q.open_access {
232 url.push_str("&openAccessPdf");
233 }
234
235 Ok(url)
236}
237
238pub fn search(base_url: &str, q: &super::SearchQuery) -> Result<Vec<Paper>, String> {
240 let url = build_search_url(base_url, q)?;
241 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
286pub 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
325pub 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
356pub 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
380pub 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
390fn 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 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 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 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 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 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 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 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 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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}