Skip to main content

update_rs/source/
github.rs

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