1use std::time::Duration;
12
13use reqwest::{Method, Response, StatusCode, header};
14use serde::de::DeserializeOwned;
15use serde_json::Value;
16use url::Url;
17use vgi_forge::{ForgeError, Result};
18
19use crate::secret::Secret;
20
21pub const API_VERSION: &str = "2022-11-28";
23
24const MAX_PAGES: usize = 50;
26
27#[derive(Clone, Copy)]
29pub(crate) enum Auth<'a> {
30 None,
32 Bearer(&'a Secret),
35}
36
37#[derive(Debug, Clone)]
38pub(crate) struct Api {
39 client: reqwest::Client,
40 pub(crate) api_base: Url,
41 pub(crate) web_base: Url,
42}
43
44impl Api {
45 pub(crate) fn new(api_base: Url, web_base: Url, timeout: Duration) -> Result<Self> {
46 let client = reqwest::Client::builder()
47 .user_agent(concat!("vgi-forge-github/", env!("CARGO_PKG_VERSION")))
48 .redirect(reqwest::redirect::Policy::none())
49 .timeout(timeout)
50 .build()
51 .map_err(|e| ForgeError::Config(format!("HTTP client: {e}")))?;
52 Ok(Api {
53 client,
54 api_base,
55 web_base,
56 })
57 }
58
59 pub(crate) fn url(&self, segments: &[&str]) -> Url {
61 join(&self.api_base, segments)
62 }
63
64 pub(crate) fn web_url(&self, segments: &[&str]) -> Url {
66 join(&self.web_base, segments)
67 }
68
69 pub(crate) async fn send(
72 &self,
73 method: Method,
74 url: Url,
75 auth: Auth<'_>,
76 body: Option<&Value>,
77 what: &str,
78 ) -> Result<Response> {
79 let resp = self.dispatch(method, url, auth, body).await?;
80 check(resp, what).await
81 }
82
83 pub(crate) async fn send_or_refusal(
87 &self,
88 method: Method,
89 url: Url,
90 auth: Auth<'_>,
91 body: Option<&Value>,
92 what: &str,
93 ) -> Result<std::result::Result<Response, Refusal>> {
94 let resp = self.dispatch(method, url, auth, body).await?;
95 let status = resp.status();
96 let refusable = status == StatusCode::UNPROCESSABLE_ENTITY
97 || (status == StatusCode::FORBIDDEN
98 && resp
99 .headers()
100 .get("x-ratelimit-remaining")
101 .and_then(|v| v.to_str().ok())
102 != Some("0")
103 && !resp.headers().contains_key(header::RETRY_AFTER));
104 if !refusable {
105 return check(resp, what).await.map(Ok);
106 }
107 let body = resp.json::<Value>().await.unwrap_or(Value::Null);
108 let message = message_of(&body);
109 let documentation_url = body
110 .get("documentation_url")
111 .and_then(Value::as_str)
112 .unwrap_or_default()
113 .to_string();
114 let error = if status == StatusCode::FORBIDDEN {
115 ForgeError::Forbidden(format!("{what}: {message}"))
116 } else {
117 ForgeError::Rejected {
118 status: status.as_u16(),
119 message: format!("{what}: {message}"),
120 }
121 };
122 Ok(Err(Refusal {
123 status: status.as_u16(),
124 message,
125 documentation_url,
126 error,
127 }))
128 }
129
130 async fn dispatch(
131 &self,
132 method: Method,
133 url: Url,
134 auth: Auth<'_>,
135 body: Option<&Value>,
136 ) -> Result<Response> {
137 let writes = matches!(method, Method::POST | Method::PUT | Method::PATCH);
138 let mut req = self
139 .client
140 .request(method, url)
141 .header(header::ACCEPT, "application/vnd.github+json")
142 .header("X-GitHub-Api-Version", API_VERSION);
143 if let Auth::Bearer(token) = auth {
144 let mut value = header::HeaderValue::try_from(format!("Bearer {}", token.expose()))
145 .map_err(|_| ForgeError::Config("credential is not a valid header value".into()))?;
146 value.set_sensitive(true);
148 req = req.header(header::AUTHORIZATION, value);
149 }
150 match body {
151 Some(body) => req = req.json(body),
152 None if writes => {
156 req = req
157 .header(header::CONTENT_LENGTH, "0")
158 .body(Vec::<u8>::new());
159 }
160 None => {}
161 }
162 req.send().await.map_err(|e| {
163 ForgeError::Unavailable(e.without_url().to_string())
165 })
166 }
167
168 pub(crate) async fn json<T: DeserializeOwned>(
170 &self,
171 method: Method,
172 url: Url,
173 auth: Auth<'_>,
174 body: Option<&Value>,
175 what: &str,
176 ) -> Result<T> {
177 let resp = self.send(method, url, auth, body, what).await?;
178 decode(resp, what).await
179 }
180
181 pub(crate) async fn oauth<T: DeserializeOwned>(&self, url: Url, body: &Value) -> Result<T> {
186 let resp = self
187 .client
188 .post(url)
189 .header(header::ACCEPT, "application/json")
190 .json(body)
191 .send()
192 .await
193 .map_err(|e| ForgeError::Unavailable(e.without_url().to_string()))?;
194 let resp = check(resp, "OAuth device flow").await?;
195 decode(resp, "OAuth device flow").await
196 }
197
198 pub(crate) async fn basic_delete(
200 &self,
201 url: Url,
202 user: &str,
203 password: &Secret,
204 body: &Value,
205 ) -> Result<()> {
206 let resp = self
207 .client
208 .delete(url)
209 .header(header::ACCEPT, "application/vnd.github+json")
210 .header("X-GitHub-Api-Version", API_VERSION)
211 .basic_auth(user, Some(password.expose()))
212 .json(body)
213 .send()
214 .await
215 .map_err(|e| ForgeError::Unavailable(e.without_url().to_string()))?;
216 check(resp, "token revocation").await.map(|_| ())
217 }
218
219 pub(crate) async fn get_opt<T: DeserializeOwned>(
221 &self,
222 url: Url,
223 auth: Auth<'_>,
224 what: &str,
225 ) -> Result<Option<T>> {
226 match self.json(Method::GET, url, auth, None, what).await {
227 Ok(v) => Ok(Some(v)),
228 Err(ForgeError::NotFound { .. }) => Ok(None),
229 Err(e) => Err(e),
230 }
231 }
232
233 pub(crate) async fn get_all<T: DeserializeOwned>(
236 &self,
237 mut url: Url,
238 auth: Auth<'_>,
239 what: &str,
240 ) -> Result<Vec<T>> {
241 url.query_pairs_mut().append_pair("per_page", "100");
242 let mut out = Vec::new();
243 for _ in 0..MAX_PAGES {
244 let resp = self
245 .send(Method::GET, url.clone(), auth, None, what)
246 .await?;
247 let next = next_link(&resp).filter(|n| n.origin() == self.api_base.origin());
248 let page: Vec<T> = decode(resp, what).await?;
249 out.extend(page);
250 match next {
251 Some(n) => url = n,
252 None => return Ok(out),
253 }
254 }
255 Err(ForgeError::Protocol(format!(
256 "{what}: more than {MAX_PAGES} pages"
257 )))
258 }
259}
260
261fn join(base: &Url, segments: &[&str]) -> Url {
268 let mut url = base.clone();
269 {
270 let mut path = url
271 .path_segments_mut()
272 .expect("API base URLs are http(s), which have paths");
273 path.pop_if_empty();
274 for s in segments {
275 path.push(s);
276 }
277 }
278 url
279}
280
281async fn decode<T: DeserializeOwned>(resp: Response, what: &str) -> Result<T> {
282 let bytes = resp
283 .bytes()
284 .await
285 .map_err(|e| ForgeError::Unavailable(e.without_url().to_string()))?;
286 serde_json::from_slice(&bytes).map_err(|e| ForgeError::Protocol(format!("{what}: {e}")))
287}
288
289async fn check(resp: Response, what: &str) -> Result<Response> {
290 let status = resp.status();
291 if status.is_success() {
292 return Ok(resp);
293 }
294 if status.is_redirection() {
295 let location = resp
296 .headers()
297 .get(header::LOCATION)
298 .and_then(|v| v.to_str().ok())
299 .unwrap_or("(no location)")
300 .to_string();
301 return Err(ForgeError::Moved {
302 what: what.to_string(),
303 location,
304 });
305 }
306
307 let headers = resp.headers().clone();
308 let message = error_message(resp).await;
309 let header_u64 = |name: &str| {
310 headers
311 .get(name)
312 .and_then(|v| v.to_str().ok())
313 .and_then(|v| v.parse::<u64>().ok())
314 };
315 let rate_limited = status == StatusCode::TOO_MANY_REQUESTS
316 || (status == StatusCode::FORBIDDEN
317 && (header_u64("x-ratelimit-remaining") == Some(0)
318 || headers.contains_key(header::RETRY_AFTER)));
319
320 Err(match status {
321 _ if rate_limited => ForgeError::RateLimited {
322 retry_after_secs: header_u64(header::RETRY_AFTER.as_str()),
323 },
324 StatusCode::UNAUTHORIZED => ForgeError::Unauthorized(format!("{what}: {message}")),
325 StatusCode::FORBIDDEN => ForgeError::Forbidden(format!("{what}: {message}")),
326 StatusCode::NOT_FOUND => ForgeError::NotFound {
327 what: what.to_string(),
328 },
329 StatusCode::CONFLICT | StatusCode::UNPROCESSABLE_ENTITY | StatusCode::BAD_REQUEST => {
330 ForgeError::Rejected {
331 status: status.as_u16(),
332 message: format!("{what}: {message}"),
333 }
334 }
335 s if s.is_server_error() => ForgeError::Unavailable(format!("{what}: {s} {message}")),
336 s => ForgeError::Protocol(format!("{what}: unexpected {s} {message}")),
337 })
338}
339
340#[derive(Debug)]
342pub(crate) struct Refusal {
343 pub(crate) status: u16,
344 pub(crate) message: String,
346 pub(crate) documentation_url: String,
348 pub(crate) error: ForgeError,
350}
351
352async fn error_message(resp: Response) -> String {
355 let Ok(body) = resp.json::<Value>().await else {
356 return "(no message)".into();
357 };
358 message_of(&body)
359}
360
361fn message_of(body: &Value) -> String {
362 let mut msg = body
363 .get("message")
364 .and_then(Value::as_str)
365 .unwrap_or("(no message)")
366 .to_string();
367 if let Some(detail) = body
368 .get("errors")
369 .and_then(Value::as_array)
370 .and_then(|e| e.first())
371 {
372 let detail = detail
373 .get("message")
374 .and_then(Value::as_str)
375 .map(str::to_string)
376 .or_else(|| detail.as_str().map(str::to_string))
377 .unwrap_or_else(|| detail.to_string());
378 msg = format!("{msg} ({detail})");
379 }
380 if msg.len() > 300 {
381 let mut end = 300;
382 while !msg.is_char_boundary(end) {
383 end -= 1;
384 }
385 msg.truncate(end);
386 msg.push('…');
387 }
388 msg
389}
390
391fn next_link(resp: &Response) -> Option<Url> {
392 let link = resp.headers().get(header::LINK)?.to_str().ok()?;
393 link.split(',').find_map(|part| {
394 let (target, params) = part.split_once(';')?;
395 if !params.split(';').any(|p| p.trim() == r#"rel="next""#) {
396 return None;
397 }
398 Url::parse(target.trim().trim_start_matches('<').trim_end_matches('>')).ok()
399 })
400}
401
402#[cfg(test)]
403mod tests {
404 use super::*;
405
406 #[test]
407 fn url_segments_are_encoded_one_by_one() {
408 let base = Url::parse("https://ghe.example/api/v3").unwrap();
409 let url = join(&base, &["repos", "acme", "a/b?c", "contents"]);
410 assert_eq!(
411 url.as_str(),
412 "https://ghe.example/api/v3/repos/acme/a%2Fb%3Fc/contents"
413 );
414 let url = join(&base, &["repos", "acme", "..", "..", "..", "x"]);
415 assert!(url.path().starts_with("/api/v3/repos"), "{url}");
416 let root = Url::parse("https://api.github.com").unwrap();
417 assert_eq!(
418 join(&root, &["user"]).as_str(),
419 "https://api.github.com/user"
420 );
421 }
422}