smriti/services/
self_replace.rs1use 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#[derive(Debug, Clone)]
26#[allow(dead_code)] pub enum InstallOutcome {
29 ReplacedRestartRequired,
32
33 InstallerLaunched { installer: String },
36
37 HandedOffToPackageManager { name: String, upgrade_cmd: String },
41
42 OpenedReleasePage { url: String },
45}
46
47#[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
82pub 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 InstallMethod::PackageManager { name, upgrade_cmd } => {
98 Ok(InstallOutcome::HandedOffToPackageManager { name, upgrade_cmd })
99 }
100
101 InstallMethod::Unknown | InstallMethod::SourceBuild => {
103 open_in_browser(&latest.html_url)?;
104 Ok(InstallOutcome::OpenedReleasePage {
105 url: latest.html_url,
106 })
107 }
108
109 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 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 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
139fn 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
169async 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 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 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
239async 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
283pub(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
306async 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 #[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(¤t_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 .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 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}