Skip to main content

smriti/services/
self_replace.rs

1//! Per-install-method update download + install logic.
2//!
3//! The update checker hands us a `LatestRelease` and the detected
4//! `InstallMethod`; this module picks the right artifact, downloads
5//! it, verifies its SHA256 against the published `SHA256SUMS`, and
6//! either self-replaces the running binary or triggers the platform
7//! installer.
8//!
9//! Critical invariant: for package-manager installs we **never**
10//! download or replace anything. The package manager owns the
11//! binary; stomping on it would confuse the system. Callers should
12//! check `InstallMethod` before calling `install_update`, but the
13//! function defensively short-circuits on those variants anyway.
14
15use std::path::{Path, PathBuf};
16
17use futures::StreamExt;
18use sha2::{Digest, Sha256};
19use tokio::io::AsyncWriteExt;
20
21use super::install_method::InstallMethod;
22use super::update_checker::{LatestRelease, ReleaseAsset};
23
24/// What happened when the user clicked "Download update".
25#[derive(Debug, Clone)]
26#[allow(dead_code)] // variant payload fields are observed via match
27                    // arms; Rust's dead-code pass can't see through.
28pub enum InstallOutcome {
29    /// The running binary was replaced in place. The UI should prompt
30    /// the user to relaunch.
31    ReplacedRestartRequired,
32
33    /// A platform installer was spawned (msiexec, `open dmg`). The app
34    /// should exit so the installer can complete. The UI prompts.
35    InstallerLaunched { installer: String },
36
37    /// The user's install method is managed externally (apt, brew, ...).
38    /// Banner should have shown the upgrade command; we refused to
39    /// touch the install.
40    HandedOffToPackageManager { name: String, upgrade_cmd: String },
41
42    /// The install method is unknown; open the release page in the
43    /// user's default browser as a fallback.
44    OpenedReleasePage { url: String },
45}
46
47/// Errors from the install pipeline. Distinguished so the UI can
48/// decide whether to show a retry affordance.
49#[derive(Debug, thiserror::Error)]
50pub enum UpdateError {
51    #[error("no matching release asset for this platform")]
52    NoAssetForPlatform,
53
54    #[error("SHA256SUMS file missing from the release")]
55    NoChecksumFile,
56
57    #[error("downloaded {got} bytes, expected {expected}")]
58    SizeMismatch { got: u64, expected: u64 },
59
60    #[error("SHA256 checksum mismatch for {asset}")]
61    ChecksumMismatch { asset: String },
62
63    #[error("SHA256SUMS does not list an entry for {asset}")]
64    ChecksumEntryMissing { asset: String },
65
66    #[error("network error while downloading: {0}")]
67    Network(String),
68
69    #[error("filesystem error: {0}")]
70    Io(#[from] std::io::Error),
71
72    #[error("failed to replace running binary: {0}")]
73    Replace(String),
74
75    #[error("failed to launch platform installer: {0}")]
76    SpawnInstaller(String),
77
78    #[error("failed to open the release page: {0}")]
79    OpenBrowser(String),
80}
81
82/// Full install flow. Dispatches to the right strategy based on the
83/// detected install method.
84///
85/// `progress` is an optional callback invoked with
86/// `(downloaded_bytes, total_bytes)` as bytes stream in, so the UI
87/// can render a progress bar. Called at most every ~256 KB.
88pub type ProgressCallback = Box<dyn FnMut(u64, u64) + Send>;
89
90pub async fn install_update(
91    method: InstallMethod,
92    latest: LatestRelease,
93    progress: Option<ProgressCallback>,
94) -> Result<InstallOutcome, UpdateError> {
95    match method {
96        // Managed externally — never download.
97        InstallMethod::PackageManager { name, upgrade_cmd } => {
98            Ok(InstallOutcome::HandedOffToPackageManager { name, upgrade_cmd })
99        }
100
101        // Unknown install → open the release page instead of guessing.
102        InstallMethod::Unknown | InstallMethod::SourceBuild => {
103            open_in_browser(&latest.html_url)?;
104            Ok(InstallOutcome::OpenedReleasePage {
105                url: latest.html_url,
106            })
107        }
108
109        // AppImage on Linux or portable zip on Windows: atomic swap.
110        InstallMethod::Portable => {
111            let asset = pick_portable_asset(&latest.assets)?.clone();
112            let downloaded = download_and_verify(&asset, &latest.assets, progress).await?;
113            replace_running_binary(&downloaded).await?;
114            Ok(InstallOutcome::ReplacedRestartRequired)
115        }
116
117        // Windows MSI: download + invoke msiexec with UAC.
118        InstallMethod::MsiInstaller => {
119            let asset = pick_msi_asset(&latest.assets)?.clone();
120            let downloaded = download_and_verify(&asset, &latest.assets, progress).await?;
121            launch_msi(&downloaded)?;
122            Ok(InstallOutcome::InstallerLaunched {
123                installer: asset.name,
124            })
125        }
126
127        // macOS .dmg: download + open so the user drags into /Applications.
128        InstallMethod::MacOsApp => {
129            let asset = pick_dmg_asset(&latest.assets)?.clone();
130            let downloaded = download_and_verify(&asset, &latest.assets, progress).await?;
131            open_dmg(&downloaded)?;
132            Ok(InstallOutcome::InstallerLaunched {
133                installer: asset.name,
134            })
135        }
136    }
137}
138
139// --- Asset selection --------------------------------------------------
140
141fn pick_portable_asset(assets: &[ReleaseAsset]) -> Result<&ReleaseAsset, UpdateError> {
142    #[cfg(target_os = "linux")]
143    let patterns: &[&str] = &[".AppImage"];
144    #[cfg(target_os = "windows")]
145    let patterns: &[&str] = &["pc-windows-msvc.zip", "windows-x64.zip"];
146    #[cfg(target_os = "macos")]
147    let patterns: &[&str] = &["apple-darwin.tar.gz"];
148
149    assets
150        .iter()
151        .find(|a| patterns.iter().any(|p| a.name.contains(p)))
152        .ok_or(UpdateError::NoAssetForPlatform)
153}
154
155fn pick_msi_asset(assets: &[ReleaseAsset]) -> Result<&ReleaseAsset, UpdateError> {
156    assets
157        .iter()
158        .find(|a| a.name.to_ascii_lowercase().ends_with(".msi"))
159        .ok_or(UpdateError::NoAssetForPlatform)
160}
161
162fn pick_dmg_asset(assets: &[ReleaseAsset]) -> Result<&ReleaseAsset, UpdateError> {
163    assets
164        .iter()
165        .find(|a| a.name.to_ascii_lowercase().ends_with(".dmg"))
166        .ok_or(UpdateError::NoAssetForPlatform)
167}
168
169// --- Download + verify -------------------------------------------------
170
171/// Download `asset` to a temp file and verify its SHA256 matches the
172/// entry in the release's `SHA256SUMS` asset. Returns the path to the
173/// verified temp file; caller is responsible for moving it into place.
174async fn download_and_verify(
175    asset: &ReleaseAsset,
176    all_assets: &[ReleaseAsset],
177    mut progress: Option<ProgressCallback>,
178) -> Result<PathBuf, UpdateError> {
179    let expected_hash = fetch_expected_hash(all_assets, &asset.name).await?;
180
181    let temp_path = std::env::temp_dir().join(format!(
182        "smriti-update-{}",
183        safe_temp_file_name(&asset.name)
184    ));
185    let client = build_client()?;
186
187    let response = client
188        .get(&asset.browser_download_url)
189        .send()
190        .await
191        .and_then(|r| r.error_for_status())
192        .map_err(|e| UpdateError::Network(e.to_string()))?;
193
194    let total = response.content_length().unwrap_or(asset.size_bytes);
195    let mut file = tokio::fs::File::create(&temp_path).await?;
196    let mut hasher = Sha256::new();
197    let mut downloaded: u64 = 0;
198    let mut last_progress_report: u64 = 0;
199
200    let mut stream = response.bytes_stream();
201    while let Some(chunk) = stream.next().await {
202        let chunk = chunk.map_err(|e| UpdateError::Network(e.to_string()))?;
203        hasher.update(&chunk);
204        file.write_all(&chunk).await?;
205        downloaded += chunk.len() as u64;
206
207        if let Some(cb) = progress.as_mut() {
208            // Report at most once per ~256 KB to keep the UI event
209            // stream from drowning in ticks.
210            if downloaded - last_progress_report > 256 * 1024 || downloaded == total {
211                cb(downloaded, total);
212                last_progress_report = downloaded;
213            }
214        }
215    }
216    file.flush().await?;
217    drop(file);
218
219    if asset.size_bytes != 0 && downloaded != asset.size_bytes {
220        return Err(UpdateError::SizeMismatch {
221            got: downloaded,
222            expected: asset.size_bytes,
223        });
224    }
225
226    let actual_hash = format!("{:x}", hasher.finalize());
227    if !actual_hash.eq_ignore_ascii_case(&expected_hash) {
228        // Clean up the tainted file so we don't accidentally install
229        // it on a retry.
230        let _ = tokio::fs::remove_file(&temp_path).await;
231        return Err(UpdateError::ChecksumMismatch {
232            asset: asset.name.clone(),
233        });
234    }
235
236    Ok(temp_path)
237}
238
239/// Download and parse the `SHA256SUMS` file, returning the expected
240/// hash for `asset_name`.
241async fn fetch_expected_hash(
242    assets: &[ReleaseAsset],
243    asset_name: &str,
244) -> Result<String, UpdateError> {
245    let sums_asset = assets
246        .iter()
247        .find(|a| a.name.eq_ignore_ascii_case("SHA256SUMS"))
248        .ok_or(UpdateError::NoChecksumFile)?;
249
250    let client = build_client()?;
251    let body = client
252        .get(&sums_asset.browser_download_url)
253        .send()
254        .await
255        .and_then(|r| r.error_for_status())
256        .map_err(|e| UpdateError::Network(e.to_string()))?
257        .text()
258        .await
259        .map_err(|e| UpdateError::Network(e.to_string()))?;
260
261    parse_sha256sums(&body, asset_name).ok_or_else(|| UpdateError::ChecksumEntryMissing {
262        asset: asset_name.to_string(),
263    })
264}
265
266fn safe_temp_file_name(name: &str) -> String {
267    let mut out = String::with_capacity(name.len().max(1));
268    for ch in name.chars() {
269        if ch.is_ascii_alphanumeric() || matches!(ch, '.' | '-' | '_') {
270            out.push(ch);
271        } else {
272            out.push('_');
273        }
274    }
275    let trimmed = out.trim_matches('.');
276    if trimmed.is_empty() {
277        "download".into()
278    } else {
279        trimmed.to_string()
280    }
281}
282
283/// Parse a single entry out of a standard `sha256sum` file. Each line
284/// looks like: `<64-hex-chars>  <filename>` or `<64-hex-chars> *<filename>`.
285pub(crate) fn parse_sha256sums(body: &str, want: &str) -> Option<String> {
286    for line in body.lines() {
287        let mut parts = line.split_whitespace();
288        let Some(hash) = parts.next() else { continue };
289        let Some(name) = parts.next() else { continue };
290        let name = name.trim_start_matches('*');
291        if name == want {
292            return Some(hash.to_lowercase());
293        }
294    }
295    None
296}
297
298fn build_client() -> Result<reqwest::Client, UpdateError> {
299    reqwest::Client::builder()
300        .timeout(std::time::Duration::from_secs(120))
301        .user_agent(format!("smriti/{}", env!("CARGO_PKG_VERSION")))
302        .build()
303        .map_err(|e| UpdateError::Network(e.to_string()))
304}
305
306// --- Binary swap / installer spawn ------------------------------------
307
308async fn replace_running_binary(new_binary: &Path) -> Result<(), UpdateError> {
309    let current = std::env::current_exe().map_err(UpdateError::Io)?;
310    let new_binary = new_binary.to_path_buf();
311    let current_clone = current.clone();
312
313    tokio::task::spawn_blocking(move || -> Result<(), UpdateError> {
314        // On Linux we also need the executable bit set on the
315        // downloaded AppImage before the move.
316        #[cfg(unix)]
317        {
318            use std::os::unix::fs::PermissionsExt;
319            let mut perms = std::fs::metadata(&new_binary)?.permissions();
320            perms.set_mode(perms.mode() | 0o755);
321            std::fs::set_permissions(&new_binary, perms)?;
322        }
323
324        self_update::Move::from_source(&new_binary)
325            .to_dest(&current_clone)
326            .map_err(|e| UpdateError::Replace(e.to_string()))
327    })
328    .await
329    .map_err(|e| UpdateError::Replace(format!("spawn_blocking panicked: {}", e)))??;
330
331    Ok(())
332}
333
334#[cfg(target_os = "windows")]
335fn launch_msi(msi: &Path) -> Result<(), UpdateError> {
336    use std::os::windows::process::CommandExt;
337    std::process::Command::new("msiexec")
338        .arg("/i")
339        .arg(msi)
340        .arg("/passive")
341        .arg("/norestart")
342        // DETACHED_PROCESS | CREATE_NEW_PROCESS_GROUP so our exit
343        // doesn't take msiexec down with us.
344        .creation_flags(0x0000_0008 | 0x0000_0200)
345        .spawn()
346        .map(|_| ())
347        .map_err(|e| UpdateError::SpawnInstaller(e.to_string()))
348}
349
350#[cfg(not(target_os = "windows"))]
351fn launch_msi(_msi: &Path) -> Result<(), UpdateError> {
352    Err(UpdateError::SpawnInstaller(
353        "MSI launch is only supported on Windows".into(),
354    ))
355}
356
357#[cfg(target_os = "macos")]
358fn open_dmg(dmg: &Path) -> Result<(), UpdateError> {
359    std::process::Command::new("open")
360        .arg("-g")
361        .arg(dmg)
362        .spawn()
363        .map(|_| ())
364        .map_err(|e| UpdateError::SpawnInstaller(e.to_string()))
365}
366
367#[cfg(not(target_os = "macos"))]
368fn open_dmg(_dmg: &Path) -> Result<(), UpdateError> {
369    Err(UpdateError::SpawnInstaller(
370        "DMG launch is only supported on macOS".into(),
371    ))
372}
373
374fn open_in_browser(url: &str) -> Result<(), UpdateError> {
375    open::that(url).map_err(|e| UpdateError::OpenBrowser(e.to_string()))
376}
377
378#[cfg(test)]
379mod tests {
380    use super::*;
381
382    #[test]
383    fn sha256sums_parse_plain_format() {
384        let body = "\
385abc123  Smriti-x86_64.AppImage
386def456  Smriti-Setup-x64.msi
387";
388        assert_eq!(
389            parse_sha256sums(body, "Smriti-x86_64.AppImage"),
390            Some("abc123".to_string())
391        );
392        assert_eq!(
393            parse_sha256sums(body, "Smriti-Setup-x64.msi"),
394            Some("def456".to_string())
395        );
396    }
397
398    #[test]
399    fn sha256sums_parse_binary_star_format() {
400        // Some sha256sum implementations use `*` to mark binary mode.
401        let body = "abc123 *Smriti-x86_64.AppImage\n";
402        assert_eq!(
403            parse_sha256sums(body, "Smriti-x86_64.AppImage"),
404            Some("abc123".to_string())
405        );
406    }
407
408    #[test]
409    fn sha256sums_returns_none_when_missing() {
410        let body = "abc123  other-file\n";
411        assert!(parse_sha256sums(body, "Smriti-x86_64.AppImage").is_none());
412    }
413
414    #[test]
415    fn sha256sums_skips_blank_lines() {
416        let body = "\n\nabc123  Smriti-x86_64.AppImage\n\n";
417        assert_eq!(
418            parse_sha256sums(body, "Smriti-x86_64.AppImage"),
419            Some("abc123".to_string())
420        );
421    }
422
423    #[test]
424    fn temp_file_name_removes_path_separators() {
425        assert_eq!(
426            safe_temp_file_name("../Smriti Setup/x64.msi"),
427            "_Smriti_Setup_x64.msi"
428        );
429        assert_eq!(safe_temp_file_name("..."), "download");
430    }
431}