1use crate::http::HttpClient;
36
37#[derive(Debug, Clone, Copy, PartialEq, Eq)]
39pub struct BaseOverride {
40 pub source: &'static str,
43 pub env: &'static str,
45 pub selects_test_client: bool,
48}
49
50const fn row(source: &'static str, env: &'static str, selects: bool) -> BaseOverride {
51 BaseOverride {
52 source,
53 env,
54 selects_test_client: selects,
55 }
56}
57
58pub const BASE_OVERRIDES: &[BaseOverride] = &[
63 row("arxiv", "DOIGET_ARXIV_BASE", true),
65 row("arxiv", "DOIGET_ARXIV_SRC_BASE", true),
66 row("crossref", "DOIGET_CROSSREF_BASE", true),
67 row("unpaywall", "DOIGET_UNPAYWALL_BASE", true),
68 row("oa-publisher", "DOIGET_OA_PUBLISHER_BASE", true),
69 row("openalex", "DOIGET_OPENALEX_BASE", true),
70 row("ar5iv", "DOIGET_AR5IV_BASE", true),
71 row("biorxiv", "DOIGET_BIORXIV_BASE", false),
73 row("inspire", "DOIGET_INSPIRE_BASE", false),
74 row("ads", "DOIGET_ADS_BASE", false),
75 row("ncbi", "DOIGET_NCBI_BASE", true),
77 row("github", "DOIGET_GITHUB_API_BASE", true),
79 row("github-raw", "DOIGET_GITHUB_RAW_BASE", true),
80 row("datacite", "DOIGET_DATACITE_BASE", false),
82 row("europe-pmc", "DOIGET_EUROPE_PMC_BASE", false),
83 row("openaire", "DOIGET_OPENAIRE_BASE", false),
84 row("hal", "DOIGET_HAL_BASE", false),
85 row("core", "DOIGET_CORE_BASE", false),
86 row("tdm-aps", "DOIGET_APS_BASE", false),
89 row("tdm-elsevier", "DOIGET_ELSEVIER_BASE", false),
90 row("tdm-springer", "DOIGET_SPRINGER_BASE", false),
91 row("tdm-ieee", "DOIGET_IEEE_BASE", false),
92];
93
94#[derive(Debug, thiserror::Error)]
96pub enum BaseOverrideError {
97 #[error("{env} is not a URL with a host: {value:?}")]
99 NotAUrl {
100 env: &'static str,
102 value: String,
104 },
105}
106
107pub fn test_client_from_env() -> Result<Option<HttpClient>, BaseOverrideError> {
120 test_client_from(|env| std::env::var(env).ok())
121}
122
123pub fn test_client_from(
129 lookup: impl Fn(&str) -> Option<String>,
130) -> Result<Option<HttpClient>, BaseOverrideError> {
131 match test_entries(lookup)? {
132 Some(entries) => {
133 let borrowed: Vec<(&str, &str)> =
134 entries.iter().map(|(s, h)| (*s, h.as_str())).collect();
135 Ok(Some(HttpClient::new_for_tests_allow_http_multi(&borrowed)))
136 }
137 None => Ok(None),
138 }
139}
140
141fn test_entries(
145 lookup: impl Fn(&str) -> Option<String>,
146) -> Result<Option<Vec<(&'static str, String)>>, BaseOverrideError> {
147 let set: Vec<(BaseOverride, String)> = BASE_OVERRIDES
148 .iter()
149 .filter_map(|o| lookup(o.env).map(|v| (*o, v)))
150 .collect();
151 if !set.iter().any(|(o, _)| o.selects_test_client) {
152 return Ok(None);
153 }
154 let mut entries: Vec<(&'static str, String)> = Vec::new();
155 for (o, value) in set {
156 if entries.iter().any(|(s, _)| *s == o.source) {
157 continue;
158 }
159 let host = url::Url::parse(&value)
160 .ok()
161 .and_then(|u| u.host_str().map(str::to_string))
162 .ok_or_else(|| BaseOverrideError::NotAUrl {
163 env: o.env,
164 value: value.clone(),
165 })?;
166 entries.push((o.source, host));
167 }
168 Ok(Some(entries))
169}
170
171#[cfg(test)]
172#[allow(clippy::expect_used, clippy::unwrap_used)]
173mod tests {
174 use super::*;
175
176 fn env<'a>(pairs: &'a [(&'a str, &'a str)]) -> impl Fn(&str) -> Option<String> + 'a {
177 move |k| {
178 pairs
179 .iter()
180 .find(|(n, _)| *n == k)
181 .map(|(_, v)| (*v).to_string())
182 }
183 }
184
185 #[test]
186 fn nothing_set_is_production() {
187 assert!(test_entries(env(&[])).unwrap().is_none());
188 }
189
190 #[test]
193 fn a_non_selecting_base_alone_stays_on_the_production_client() {
194 for o in BASE_OVERRIDES.iter().filter(|o| !o.selects_test_client) {
195 let pairs = [(o.env, "https://proxy.example.edu")];
196 assert!(test_entries(env(&pairs)).unwrap().is_none(), "{}", o.env);
197 }
198 }
199
200 #[test]
203 fn test_mode_registers_every_overridden_source_and_nothing_else() {
204 let pairs = [
205 ("DOIGET_CROSSREF_BASE", "http://127.0.0.1:9001"),
206 ("DOIGET_DATACITE_BASE", "http://127.0.0.1:9002"),
207 ("DOIGET_EUROPE_PMC_BASE", "http://127.0.0.1:9003"),
208 ("DOIGET_APS_BASE", "http://127.0.0.1:9004"),
209 ];
210 let entries = test_entries(env(&pairs)).unwrap().expect("test mode");
211 let keys: Vec<&str> = entries.iter().map(|(s, _)| *s).collect();
212 assert_eq!(keys, vec!["crossref", "datacite", "europe-pmc", "tdm-aps"]);
213 assert!(entries.iter().all(|(_, h)| h == "127.0.0.1"));
214 }
215
216 #[test]
217 fn the_source_bundle_base_stands_in_for_arxiv_only_when_it_is_unset() {
218 let src_only = [("DOIGET_ARXIV_SRC_BASE", "http://src.test:1")];
219 let entries = test_entries(env(&src_only)).unwrap().expect("test mode");
220 assert_eq!(entries, vec![("arxiv", "src.test".to_string())]);
221 let both = [
222 ("DOIGET_ARXIV_BASE", "http://api.test:1"),
223 ("DOIGET_ARXIV_SRC_BASE", "http://src.test:1"),
224 ];
225 let entries = test_entries(env(&both)).unwrap().expect("test mode");
226 assert_eq!(entries, vec![("arxiv", "api.test".to_string())]);
227 }
228
229 #[test]
230 fn a_malformed_base_in_test_mode_names_its_variable() {
231 let pairs = [
232 ("DOIGET_CROSSREF_BASE", "http://127.0.0.1:1"),
233 ("DOIGET_HAL_BASE", "not a url"),
234 ];
235 let err = test_entries(env(&pairs)).unwrap_err();
236 assert!(
237 err.to_string().starts_with("DOIGET_HAL_BASE is not a URL"),
238 "{err}"
239 );
240 }
241
242 #[test]
248 fn every_base_override_read_by_a_source_has_a_row() {
249 let sources = [
250 include_str!("orchestrator.rs"),
251 include_str!("paper_tex_source.rs"),
252 include_str!("paper_text.rs"),
253 include_str!("discovery.rs"),
254 include_str!("citation_graph.rs"),
255 ]
256 .concat();
257 let mut read: Vec<&str> = sources
258 .match_indices("\"DOIGET_")
259 .filter_map(|(i, _)| {
260 let rest = &sources[i + 1..];
261 let name = &rest[..rest.find('"')?];
262 name.ends_with("_BASE").then_some(name)
263 })
264 .collect();
265 read.sort_unstable();
266 read.dedup();
267 assert!(read.len() >= 10, "scan found too few variables: {read:?}");
268 let missing: Vec<&str> = read
269 .into_iter()
270 .filter(|v| !BASE_OVERRIDES.iter().any(|o| o.env == *v))
271 .collect();
272 assert!(
273 missing.is_empty(),
274 "DOIGET_*_BASE without a BASE_OVERRIDES row: {missing:?}"
275 );
276 }
277}