Skip to main content

update_rs/source/
github.rs

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
11/// The `User-Agent` header sent with every request, built from this crate's
12/// version at compile time. GitHub requires a `User-Agent` on all API requests.
13const USER_AGENT: &str = concat!("SierraSoftworks/update-rs v", env!("CARGO_PKG_VERSION"));
14
15/// A [`Source`] which lists and downloads releases from a GitHub repository's
16/// [releases API](https://docs.github.com/en/rest/releases).
17///
18/// The asset to download is selected by matching a **glob pattern** against each
19/// release's asset file names, so your project can name its release assets
20/// however it likes — there is no required naming scheme. The pattern is the
21/// second argument to [`new`](GitHubSource::new):
22///
23/// ```
24/// use std::env::consts::{ARCH, EXE_SUFFIX, OS};
25/// use update_rs::GitHubSource;
26///
27/// let source = GitHubSource::new(
28///     "sierrasoftworks/git-tool",
29///     format!("git-tool-{OS}-{ARCH}{EXE_SUFFIX}"),
30/// )
31/// .with_release_tag_prefix("v");
32/// ```
33///
34/// The [`naming`](crate::naming) helpers build common patterns for you, e.g.
35/// [`naming::go`](crate::naming::go) (`git-tool-linux-amd64`) or
36/// [`naming::rust`](crate::naming::rust) (`git-tool-x86_64-unknown-linux-gnu`).
37///
38/// The pattern must match the asset's **whole** file name (it is anchored at
39/// both ends), so an exact name won't accidentally select a `.sha256` checksum
40/// or `.sig` sidecar. It supports `*` (any sequence) and `?` (single character)
41/// wildcards; every other character matches literally. `release_tag_prefix` is
42/// stripped from each Git tag before it is parsed as a
43/// [semantic version](semver::Version), and tags that don't parse afterwards are
44/// ignored.
45///
46/// When GitHub reports a SHA-256 [digest](https://docs.github.com/en/rest/releases/assets)
47/// for an asset, the downloaded bytes are verified against it before the update
48/// proceeds, so a corrupted or tampered download is rejected.
49pub 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    /// Create a source for the `owner/name` `repo` (e.g.
61    /// `"sierrasoftworks/git-tool"`), selecting the asset to download with the
62    /// glob `asset_pattern` (e.g. `"git-tool-linux-amd64"` or `"*-linux-amd64"`).
63    ///
64    /// See the [`naming`](crate::naming) module for helpers that build a pattern
65    /// for the current platform.
66    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    /// Strip `prefix` from each release's Git tag before parsing it as a
79    /// [semantic version](semver::Version) (e.g. `"v"` for `vX.Y.Z` tags).
80    pub fn with_release_tag_prefix(mut self, prefix: &str) -> Self {
81        self.release_tag_prefix = prefix.to_string();
82        self
83    }
84
85    /// Override the GitHub web (`web`) and API (`api`) endpoints. This is
86    /// primarily useful for pointing the source at a GitHub Enterprise instance
87    /// or a mock server in tests. Trailing slashes are trimmed.
88    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    /// Select the first release asset whose name matches the configured glob
215    /// pattern, capturing its SHA-256 digest if GitHub reported one.
216    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
303/// Parse a GitHub asset `digest` field (e.g. `"sha256:abc..."`) into its
304/// lowercase hex SHA-256, returning `None` for absent or unsupported algorithms.
305fn parse_sha256_digest(digest: &str) -> Option<String> {
306    digest
307        .strip_prefix("sha256:")
308        .map(|hex| hex.trim().to_ascii_lowercase())
309}
310
311/// Encode bytes as lowercase hexadecimal.
312fn 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        // Only the matching asset becomes the variant.
399        assert!(release.get_variant().is_some());
400        let variant = release.get_variant().unwrap();
401        assert_eq!(variant.name, "update-linux-amd64");
402        // No digest in this fixture.
403        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        // A glob that matches a single asset regardless of host platform.
416        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        // An exact pattern must select the binary, not its checksum/signature
462        // sidecars, even when they are listed first.
463        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}