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