1use super::Source;
2use crate::{Error, Release, ReleaseVariant, glob};
3use futures_util::StreamExt;
4use human_errors::ResultExt;
5use reqwest::StatusCode;
6use serde::Deserialize;
7use sha2::{Digest, Sha256};
8use std::io::Write;
9use tracing::debug;
10
11const USER_AGENT: &str = concat!("SierraSoftworks/update-rs v", env!("CARGO_PKG_VERSION"));
14
15pub struct GitHubSource {
50 github_endpoint: String,
51 github_api: String,
52 repo: String,
53 asset_pattern: String,
54 release_tag_prefix: String,
55
56 client: reqwest::Client,
57}
58
59impl GitHubSource {
60 pub fn new(repo: impl Into<String>, asset_pattern: impl Into<String>) -> Self {
67 Self {
68 github_endpoint: "https://github.com".to_string(),
69 github_api: "https://api.github.com".to_string(),
70 repo: repo.into(),
71 asset_pattern: asset_pattern.into(),
72 release_tag_prefix: String::new(),
73
74 client: reqwest::Client::new(),
75 }
76 }
77
78 pub fn with_release_tag_prefix(mut self, prefix: &str) -> Self {
81 self.release_tag_prefix = prefix.to_string();
82 self
83 }
84
85 pub fn with_github_endpoints(mut self, web: &str, api: &str) -> Self {
89 self.github_endpoint = web.trim_end_matches('/').to_string();
90 self.github_api = api.trim_end_matches('/').to_string();
91 self
92 }
93}
94
95impl Default for GitHubSource {
96 fn default() -> Self {
97 GitHubSource::new("", "*")
98 }
99}
100
101impl std::fmt::Debug for GitHubSource {
102 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
103 write!(f, "GitHub - {} ({})", &self.repo, &self.asset_pattern)
104 }
105}
106
107#[async_trait::async_trait]
108impl Source for GitHubSource {
109 async fn get_releases(&self) -> Result<Vec<Release>, Error> {
110 let uri = format!("{}/repos/{}/releases", self.github_api, self.repo);
111 debug!("Making GET request to {} to check for new releases.", uri);
112
113 let resp = self.get(&uri).await?;
114 debug!(
115 "Received HTTP {} from GitHub when requesting releases.",
116 resp.status()
117 );
118
119 match resp.status() {
120 StatusCode::OK => {
121 let releases: Vec<GitHubRelease> = resp.json().await.wrap_system_err(
122 "Unable to parse the response from the GitHub releases API.",
123 &["The GitHub API may be unavailable or may have changed in an incompatible way; please report this issue if it persists."],
124 )?;
125
126 debug!("Received {} releases from GitHub.", releases.len());
127 Ok(self.get_releases_from_response(releases))
128 }
129 StatusCode::NOT_FOUND => Err(human_errors::user(
130 "GitHub returned a 404 Not Found when listing the releases for this repository.",
131 &[
132 "Check that the repository exists and is public, and that the update manager is configured with the correct 'owner/name' repository identifier.",
133 ],
134 )),
135 StatusCode::TOO_MANY_REQUESTS | StatusCode::FORBIDDEN => Err(human_errors::user(
136 "GitHub has rate limited requests from your IP address.",
137 &["Please wait until GitHub removes this rate limit before trying again."],
138 )),
139 status => {
140 let body = resp.text().await.unwrap_or_default();
141 Err(human_errors::wrap_system(
142 body,
143 format!(
144 "Received an HTTP {status} response from GitHub when listing the available releases."
145 ),
146 &[
147 "Please read the error message below and decide if there is something you can do to fix the problem, or report the issue.",
148 ],
149 ))
150 }
151 }
152 }
153
154 async fn get_binary<W: Write + Send>(
155 &self,
156 release: &Release,
157 variant: &ReleaseVariant,
158 into: &mut W,
159 ) -> Result<(), Error> {
160 let uri = format!(
161 "{}/{}/releases/download/{}/{}",
162 self.github_endpoint, self.repo, release.id, variant.name
163 );
164
165 self.download_to_file(&uri, variant.sha256.as_deref(), into)
166 .await
167 }
168}
169
170impl GitHubSource {
171 async fn get(&self, uri: &str) -> Result<reqwest::Response, Error> {
172 self.client
173 .get(uri)
174 .header("User-Agent", USER_AGENT)
175 .send()
176 .await
177 .wrap_system_err(
178 format!("Failed to make a request to '{uri}'."),
179 &["Check your network connection and try again, or report the issue if it persists."],
180 )
181 }
182
183 fn get_releases_from_response(&self, releases: Vec<GitHubRelease>) -> Vec<Release> {
184 let mut output: Vec<Release> = Vec::with_capacity(releases.len());
185
186 for r in releases {
187 if !r.tag_name.starts_with(&self.release_tag_prefix) {
188 continue;
189 }
190
191 match r.tag_name[self.release_tag_prefix.len()..].parse() {
192 Ok(version) => {
193 debug!("Found release '{}'.", r.tag_name);
194 output.push(Release {
195 id: r.tag_name.clone(),
196 changelog: r.body.clone(),
197 version,
198 prerelease: r.prerelease,
199 variant: self.get_variant_from_response(&r),
200 })
201 }
202 Err(_) => {
203 debug!(
204 "Skipping release '{}' because it is not a valid SemVer version (adjust the release tag prefix to fix this).",
205 &r.tag_name
206 );
207 }
208 }
209 }
210
211 output
212 }
213
214 fn get_variant_from_response(&self, release: &GitHubRelease) -> Option<ReleaseVariant> {
217 release
218 .assets
219 .iter()
220 .find(|a| glob::matches(&self.asset_pattern, &a.name))
221 .map(|a| ReleaseVariant {
222 name: a.name.clone(),
223 sha256: a.digest.as_deref().and_then(parse_sha256_digest),
224 })
225 }
226
227 async fn download_to_file<W: Write + Send>(
228 &self,
229 uri: &str,
230 expected_sha256: Option<&str>,
231 into: &mut W,
232 ) -> Result<(), Error> {
233 let resp = self.get(uri).await?;
234
235 match resp.status() {
236 StatusCode::OK => {
237 let mut hasher = Sha256::new();
238 let mut stream = resp.bytes_stream();
239
240 while let Some(chunk) = stream.next().await {
241 let chunk = chunk.wrap_user_err(
242 format!("Failed to download the update from '{uri}'."),
243 &["Check your network connection and try again, or report the issue if it persists."],
244 )?;
245 hasher.update(&chunk);
246 into.write_all(&chunk).wrap_user_err(
247 format!("Could not write data downloaded from '{uri}' to disk due to an OS-level error."),
248 &["Check that this tool has permission to create and write to this file and that the parent directory exists."],
249 )?;
250 }
251
252 match expected_sha256 {
253 Some(expected) => {
254 let actual = to_hex(hasher.finalize().as_slice());
255 if actual.eq_ignore_ascii_case(expected) {
256 debug!("Verified the downloaded update against its SHA-256 digest.");
257 } else {
258 return Err(human_errors::user(
259 format!(
260 "The update downloaded from '{uri}' failed its integrity check (expected SHA-256 {expected}, got {actual})."
261 ),
262 &[
263 "The download may have been corrupted in transit or tampered with. Please try the update again, and report the issue if it keeps happening.",
264 ],
265 ));
266 }
267 }
268 None => {
269 debug!(
270 "No SHA-256 digest was reported for this asset; skipping the integrity check."
271 );
272 }
273 }
274
275 Ok(())
276 }
277 StatusCode::NOT_FOUND => Err(human_errors::user(
278 format!("GitHub returned a 404 Not Found when downloading '{uri}'."),
279 &[
280 "This release variant may not be available for your platform, or the release may have been removed.",
281 ],
282 )),
283 StatusCode::TOO_MANY_REQUESTS | StatusCode::FORBIDDEN => Err(human_errors::user(
284 "GitHub has rate limited requests from your IP address.",
285 &["Please wait until GitHub removes this rate limit before trying again."],
286 )),
287 status => {
288 let body = resp.text().await.unwrap_or_default();
289 Err(human_errors::wrap_system(
290 body,
291 format!(
292 "Received an HTTP {status} response from GitHub when downloading the update ({uri})."
293 ),
294 &[
295 "Please read the error message below and decide if there is something you can do to fix the problem, or report the issue.",
296 ],
297 ))
298 }
299 }
300 }
301}
302
303fn parse_sha256_digest(digest: &str) -> Option<String> {
306 digest
307 .strip_prefix("sha256:")
308 .map(|hex| hex.trim().to_ascii_lowercase())
309}
310
311fn to_hex(bytes: &[u8]) -> String {
313 use std::fmt::Write;
314 let mut s = String::with_capacity(bytes.len() * 2);
315 for b in bytes {
316 let _ = write!(s, "{b:02x}");
317 }
318 s
319}
320
321#[derive(Debug, Deserialize)]
322struct GitHubRelease {
323 #[allow(dead_code)]
324 pub name: String,
325 pub tag_name: String,
326 pub body: String,
327 pub prerelease: bool,
328 pub assets: Vec<GitHubAsset>,
329}
330
331#[derive(Debug, Deserialize)]
332struct GitHubAsset {
333 pub name: String,
334 #[serde(default)]
335 pub digest: Option<String>,
336}
337
338#[cfg(test)]
339mod tests {
340 use super::*;
341 use std::sync::{Arc, Mutex};
342 use wiremock::matchers::{method, path};
343 use wiremock::{Mock, MockServer, ResponseTemplate};
344
345 const RELEASES_JSON: &str = r#"[
346 {
347 "name": "Version 2.0.0",
348 "tag_name": "v2.0.0",
349 "body": "Example Release",
350 "prerelease": false,
351 "assets": [
352 { "name": "update-windows-amd64.exe" },
353 { "name": "update-linux-amd64" },
354 { "name": "update-darwin-amd64" }
355 ]
356 }
357 ]"#;
358
359 fn source_for(server: &MockServer, pattern: &str) -> GitHubSource {
360 GitHubSource::new("sierrasoftworks/update-rs", pattern)
361 .with_github_endpoints(&server.uri(), &server.uri())
362 .with_release_tag_prefix("v")
363 }
364
365 fn sha256_hex(data: &[u8]) -> String {
366 to_hex(Sha256::digest(data).as_slice())
367 }
368
369 fn releases_json(asset: &str, digest: Option<&str>) -> String {
370 let digest_field = match digest {
371 Some(d) => format!(r#", "digest": "{d}""#),
372 None => String::new(),
373 };
374 format!(
375 r#"[{{"name":"Version 2.0.0","tag_name":"v2.0.0","body":"Example Release","prerelease":false,"assets":[{{"name":"{asset}"{digest_field}}}]}}]"#
376 )
377 }
378
379 #[tokio::test]
380 async fn test_get_releases_selects_matching_asset() {
381 let server = MockServer::start().await;
382 Mock::given(method("GET"))
383 .and(path("/repos/sierrasoftworks/update-rs/releases"))
384 .respond_with(ResponseTemplate::new(200).set_body_string(RELEASES_JSON))
385 .mount(&server)
386 .await;
387
388 let source = source_for(&server, "update-linux-amd64");
389 let releases = source.get_releases().await.unwrap();
390
391 assert_eq!(releases.len(), 1);
392 let release = &releases[0];
393 assert_eq!(release.id, "v2.0.0");
394 assert_eq!(release.version.to_string(), "2.0.0");
395 assert!(!release.prerelease);
396 assert_ne!(release.changelog, "");
397
398 assert!(release.get_variant().is_some());
400 let variant = release.get_variant().unwrap();
401 assert_eq!(variant.name, "update-linux-amd64");
402 assert!(variant.sha256.is_none());
404 }
405
406 #[tokio::test]
407 async fn test_get_releases_glob_pattern() {
408 let server = MockServer::start().await;
409 Mock::given(method("GET"))
410 .and(path("/repos/sierrasoftworks/update-rs/releases"))
411 .respond_with(ResponseTemplate::new(200).set_body_string(RELEASES_JSON))
412 .mount(&server)
413 .await;
414
415 let source = source_for(&server, "*-windows-amd64.exe");
417 let releases = source.get_releases().await.unwrap();
418
419 assert_eq!(
420 releases[0].get_variant().unwrap().name,
421 "update-windows-amd64.exe"
422 );
423 }
424
425 #[tokio::test]
426 async fn test_get_releases_no_match() {
427 let server = MockServer::start().await;
428 Mock::given(method("GET"))
429 .and(path("/repos/sierrasoftworks/update-rs/releases"))
430 .respond_with(ResponseTemplate::new(200).set_body_string(RELEASES_JSON))
431 .mount(&server)
432 .await;
433
434 let source = source_for(&server, "update-freebsd-amd64");
435 let releases = source.get_releases().await.unwrap();
436
437 assert_eq!(releases.len(), 1);
438 assert!(releases[0].get_variant().is_none());
439 }
440
441 #[tokio::test]
442 async fn test_get_releases_ignores_sidecar_files() {
443 let server = MockServer::start().await;
444 let body = r#"[{
445 "name": "Version 2.0.0",
446 "tag_name": "v2.0.0",
447 "body": "Example Release",
448 "prerelease": false,
449 "assets": [
450 { "name": "update-linux-amd64.sha256" },
451 { "name": "update-linux-amd64.sig" },
452 { "name": "update-linux-amd64" }
453 ]
454 }]"#;
455 Mock::given(method("GET"))
456 .and(path("/repos/sierrasoftworks/update-rs/releases"))
457 .respond_with(ResponseTemplate::new(200).set_body_string(body))
458 .mount(&server)
459 .await;
460
461 let source = source_for(&server, "update-linux-amd64");
464 let releases = source.get_releases().await.unwrap();
465
466 assert_eq!(
467 releases[0].get_variant().unwrap().name,
468 "update-linux-amd64"
469 );
470 }
471
472 #[tokio::test]
473 async fn test_download() {
474 let server = MockServer::start().await;
475 Mock::given(method("GET"))
476 .and(path("/repos/sierrasoftworks/update-rs/releases"))
477 .respond_with(ResponseTemplate::new(200).set_body_string(RELEASES_JSON))
478 .mount(&server)
479 .await;
480 Mock::given(method("GET"))
481 .and(path(
482 "/sierrasoftworks/update-rs/releases/download/v2.0.0/update-linux-amd64",
483 ))
484 .respond_with(ResponseTemplate::new(200).set_body_string("example update content"))
485 .mount(&server)
486 .await;
487
488 let source = source_for(&server, "update-linux-amd64");
489 let releases = source.get_releases().await.unwrap();
490 let latest = Release::get_latest(releases.iter()).unwrap();
491 let variant = latest.get_variant().unwrap();
492
493 let mut target = Sink::new();
494 source
495 .get_binary(latest, variant, &mut target)
496 .await
497 .unwrap();
498
499 assert!(target.len() > 0);
500 }
501
502 #[tokio::test]
503 async fn test_download_verifies_matching_sha256() {
504 let body = "example update content";
505 let server = MockServer::start().await;
506 Mock::given(method("GET"))
507 .and(path("/repos/sierrasoftworks/update-rs/releases"))
508 .respond_with(ResponseTemplate::new(200).set_body_string(releases_json(
509 "update-linux-amd64",
510 Some(&format!("sha256:{}", sha256_hex(body.as_bytes()))),
511 )))
512 .mount(&server)
513 .await;
514 Mock::given(method("GET"))
515 .and(path(
516 "/sierrasoftworks/update-rs/releases/download/v2.0.0/update-linux-amd64",
517 ))
518 .respond_with(ResponseTemplate::new(200).set_body_string(body))
519 .mount(&server)
520 .await;
521
522 let source = source_for(&server, "update-linux-amd64");
523 let releases = source.get_releases().await.unwrap();
524 let latest = Release::get_latest(releases.iter()).unwrap();
525 let variant = latest.get_variant().unwrap();
526 assert!(
527 variant.sha256.is_some(),
528 "the digest should have been parsed from the API response"
529 );
530
531 let mut target = Sink::new();
532 source
533 .get_binary(latest, variant, &mut target)
534 .await
535 .expect("a matching digest should pass verification");
536 assert!(target.len() > 0);
537 }
538
539 #[tokio::test]
540 async fn test_download_rejects_bad_sha256() {
541 let server = MockServer::start().await;
542 Mock::given(method("GET"))
543 .and(path("/repos/sierrasoftworks/update-rs/releases"))
544 .respond_with(ResponseTemplate::new(200).set_body_string(releases_json(
545 "update-linux-amd64",
546 Some(&format!("sha256:{}", "0".repeat(64))),
547 )))
548 .mount(&server)
549 .await;
550 Mock::given(method("GET"))
551 .and(path(
552 "/sierrasoftworks/update-rs/releases/download/v2.0.0/update-linux-amd64",
553 ))
554 .respond_with(ResponseTemplate::new(200).set_body_string("the actual bytes"))
555 .mount(&server)
556 .await;
557
558 let source = source_for(&server, "update-linux-amd64");
559 let releases = source.get_releases().await.unwrap();
560 let latest = Release::get_latest(releases.iter()).unwrap();
561 let variant = latest.get_variant().unwrap();
562
563 let mut target = Sink::new();
564 let err = source
565 .get_binary(latest, variant, &mut target)
566 .await
567 .expect_err("a mismatched digest must fail the update");
568 assert!(
569 err.to_string().contains("integrity check"),
570 "unexpected error: {err}"
571 );
572 }
573
574 #[test]
575 fn test_parse_sha256_digest() {
576 assert_eq!(parse_sha256_digest("sha256:ABCdef"), Some("abcdef".into()));
577 assert_eq!(parse_sha256_digest("sha512:abc"), None);
578 assert_eq!(parse_sha256_digest(""), None);
579 }
580
581 struct Sink {
582 length: Arc<Mutex<usize>>,
583 }
584
585 impl Sink {
586 fn new() -> Self {
587 Self {
588 length: Arc::new(Mutex::new(0)),
589 }
590 }
591
592 fn len(&self) -> usize {
593 *self.length.lock().unwrap()
594 }
595 }
596
597 impl Write for Sink {
598 fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
599 *self.length.lock().unwrap() += buf.len();
600 Ok(buf.len())
601 }
602
603 fn flush(&mut self) -> std::io::Result<()> {
604 Ok(())
605 }
606 }
607}