1use std::{cmp::Ordering, time::Duration};
4
5use reqwest::Client;
6use semver::Version;
7use serde::Deserialize;
8use thiserror::Error;
9
10use super::user_agent::get_pi_user_agent;
11
12pub const LATEST_VERSION_URL: &str = "https://pi.dev/api/latest-version";
14pub const DEFAULT_VERSION_CHECK_TIMEOUT_MS: u64 = 10_000;
16pub const ENV_SKIP_VERSION_CHECK: &str = "PI_SKIP_VERSION_CHECK";
18pub const ENV_OFFLINE: &str = "PI_OFFLINE";
20
21#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
23pub enum ReleaseChannel {
24 #[default]
26 Stable,
27 Beta,
29}
30
31#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq)]
33#[serde(rename_all = "camelCase")]
34pub struct LatestPiRelease {
35 pub version: String,
37 #[serde(default)]
39 pub package_name: Option<String>,
40 #[serde(default)]
42 pub note: Option<String>,
43}
44
45#[derive(Debug, Error)]
47pub enum VersionCheckError {
48 #[error("version checks are disabled while offline")]
50 Offline,
51 #[error("version endpoint request failed: {0}")]
53 Http(#[from] reqwest::Error),
54 #[error("version endpoint returned malformed release data")]
56 Malformed,
57}
58
59#[derive(Clone, Debug, Deserialize)]
60#[serde(rename_all = "camelCase")]
61struct ReleaseEnvelope {
62 #[serde(default)]
63 version: Option<String>,
64 #[serde(default)]
65 package_name: Option<String>,
66 #[serde(default)]
67 note: Option<String>,
68 #[serde(default)]
69 stable: Option<LatestPiRelease>,
70 #[serde(default)]
71 beta: Option<LatestPiRelease>,
72}
73
74#[must_use]
79pub fn compare_package_versions(left: &str, right: &str) -> Option<Ordering> {
80 Some(parse_version(left)?.cmp(&parse_version(right)?))
81}
82
83#[must_use]
85pub fn is_newer_package_version(candidate: &str, current: &str) -> bool {
86 compare_package_versions(candidate, current)
87 .map_or_else(|| candidate.trim() != current.trim(), Ordering::is_gt)
88}
89
90fn parse_version(value: &str) -> Option<Version> {
91 let trimmed = value.trim();
92 Version::parse(trimmed.strip_prefix('v').unwrap_or(trimmed)).ok()
93}
94
95#[must_use]
97pub fn should_skip_version_check(skip: Option<&str>, offline: Option<&str>) -> bool {
98 skip.is_some_and(|value| !value.is_empty()) || offline.is_some_and(|value| !value.is_empty())
99}
100
101pub async fn get_latest_pi_release_from(
109 client: &Client,
110 endpoint: &str,
111 current_version: &str,
112 timeout: Duration,
113 channel: ReleaseChannel,
114) -> Result<LatestPiRelease, VersionCheckError> {
115 let response = client
116 .get(endpoint)
117 .header("User-Agent", get_pi_user_agent(current_version))
118 .header("accept", "application/json")
119 .timeout(timeout)
120 .send()
121 .await?
122 .error_for_status()?;
123 let envelope = response.json::<ReleaseEnvelope>().await?;
124 resolve_release(envelope, channel).ok_or(VersionCheckError::Malformed)
125}
126
127fn resolve_release(envelope: ReleaseEnvelope, channel: ReleaseChannel) -> Option<LatestPiRelease> {
128 let release = match channel {
129 ReleaseChannel::Stable => envelope.stable,
130 ReleaseChannel::Beta => envelope.beta.or(envelope.stable),
131 }
132 .or_else(|| {
133 envelope.version.map(|version| LatestPiRelease {
134 version,
135 package_name: envelope.package_name,
136 note: envelope.note,
137 })
138 })?;
139
140 let version = release.version.trim();
141 if version.is_empty() {
142 return None;
143 }
144 Some(LatestPiRelease {
145 version: version.to_owned(),
146 package_name: release
147 .package_name
148 .map(|value| value.trim().to_owned())
149 .filter(|value| !value.is_empty()),
150 note: release
151 .note
152 .map(|value| value.trim().to_owned())
153 .filter(|value| !value.is_empty()),
154 })
155}
156
157pub async fn get_latest_pi_release(
165 current_version: &str,
166 channel: ReleaseChannel,
167) -> Result<LatestPiRelease, VersionCheckError> {
168 if should_skip_version_check(
169 std::env::var(ENV_SKIP_VERSION_CHECK).ok().as_deref(),
170 std::env::var(ENV_OFFLINE).ok().as_deref(),
171 ) {
172 return Err(VersionCheckError::Offline);
173 }
174 get_latest_pi_release_from(
175 &Client::new(),
176 LATEST_VERSION_URL,
177 current_version,
178 Duration::from_millis(DEFAULT_VERSION_CHECK_TIMEOUT_MS),
179 channel,
180 )
181 .await
182}
183
184pub async fn check_for_new_pi_version(current_version: &str) -> Option<LatestPiRelease> {
186 let release = get_latest_pi_release(current_version, ReleaseChannel::Stable)
187 .await
188 .ok()?;
189 is_newer_package_version(&release.version, current_version).then_some(release)
190}
191
192pub async fn check_for_new_pi_version_from(
194 client: &Client,
195 endpoint: &str,
196 current_version: &str,
197 timeout: Duration,
198 channel: ReleaseChannel,
199 offline: bool,
200) -> Option<LatestPiRelease> {
201 if offline {
202 return None;
203 }
204 let release = get_latest_pi_release_from(client, endpoint, current_version, timeout, channel)
205 .await
206 .ok()?;
207 is_newer_package_version(&release.version, current_version).then_some(release)
208}
209
210#[cfg(test)]
211mod tests {
212 use std::{
213 io::{Read, Write},
214 net::TcpListener,
215 thread,
216 time::Duration,
217 };
218
219 use super::*;
220
221 #[test]
222 fn semantic_order_covers_prereleases_and_malformed_fallback() {
223 assert_eq!(
224 compare_package_versions("1.2.3", "1.2.3"),
225 Some(Ordering::Equal)
226 );
227 assert_eq!(
228 compare_package_versions("v1.2.4", "1.2.3"),
229 Some(Ordering::Greater)
230 );
231 assert_eq!(
232 compare_package_versions("2.0.0-beta.1", "2.0.0"),
233 Some(Ordering::Less)
234 );
235 assert_eq!(compare_package_versions("1.2", "1.2.0"), None);
236 assert!(is_newer_package_version("malformed-a", "malformed-b"));
237 assert!(!is_newer_package_version(" malformed ", "malformed"));
238 }
239
240 #[test]
241 fn semver_handles_v_prefix_whitespace_and_prerelease_precedence() {
242 assert_eq!(
244 compare_package_versions("v1.0.0", "1.0.0"),
245 Some(Ordering::Equal)
246 );
247 assert_eq!(
249 compare_package_versions(" 1.0.0 ", "1.0.0"),
250 Some(Ordering::Equal)
251 );
252 assert_eq!(
254 compare_package_versions("1.0.0-alpha.1", "1.0.0-beta.1"),
255 Some(Ordering::Less)
256 );
257 assert_eq!(
258 compare_package_versions("1.0.0-rc.1", "1.0.0-alpha.1"),
259 Some(Ordering::Greater)
260 );
261 assert_eq!(
263 compare_package_versions("1.0.0-rc.1", "1.0.0-alpha.1"),
264 Some(Ordering::Greater)
265 );
266 }
267
268 #[test]
269 fn is_newer_package_version_uses_semver_then_string_fallback() {
270 assert!(is_newer_package_version("1.2.4", "1.2.3"));
272 assert!(!is_newer_package_version("1.0.0", "1.0.0"));
274 assert!(!is_newer_package_version("0.9.0", "1.0.0"));
276 assert!(!is_newer_package_version("1.0.0-beta.1", "1.0.0"));
278 assert!(is_newer_package_version("zzz", "aaa"));
280 assert!(!is_newer_package_version("aaa", "aaa"));
281 assert!(!is_newer_package_version(" aaa ", "aaa"));
283 }
284
285 #[test]
286 fn should_skip_version_check_respects_both_env_switches() {
287 assert!(!should_skip_version_check(None, None));
289 assert!(!should_skip_version_check(Some(""), Some("")));
291 assert!(should_skip_version_check(Some("1"), None));
293 assert!(should_skip_version_check(Some("false"), None));
294 assert!(should_skip_version_check(None, Some("1")));
296 assert!(should_skip_version_check(None, Some("true")));
297 assert!(should_skip_version_check(Some("1"), Some("1")));
299 }
300
301 #[test]
302 fn resolve_release_stable_uses_stable_field() -> Result<(), serde_json::Error> {
303 let envelope: ReleaseEnvelope = serde_json::from_value(serde_json::json!({
304 "stable": {"version": " 2.0.0 ", "packageName": " pi-new "},
305 "beta": {"version": "2.1.0-beta.1"}
306 }))?;
307 let result = resolve_release(envelope, ReleaseChannel::Stable);
308 assert!(result.is_some());
309 let release = result.unwrap_or_default();
310 assert_eq!(release.version, "2.0.0");
311 assert_eq!(release.package_name, Some("pi-new".to_owned()));
312 assert!(release.note.is_none());
313 Ok(())
314 }
315
316 #[test]
317 fn resolve_release_beta_falls_back_to_stable() -> Result<(), serde_json::Error> {
318 let envelope: ReleaseEnvelope = serde_json::from_value(serde_json::json!({
320 "stable": {"version": "2.0.0"}
321 }))?;
322 let release = resolve_release(envelope.clone(), ReleaseChannel::Beta);
323 assert_eq!(release.map(|r| r.version), Some("2.0.0".to_owned()));
324
325 let envelope: ReleaseEnvelope = serde_json::from_value(serde_json::json!({
327 "stable": {"version": "2.0.0"},
328 "beta": {"version": "2.1.0-beta.1"}
329 }))?;
330 let result = resolve_release(envelope, ReleaseChannel::Beta);
331 assert!(result.is_some());
332 let release = result.unwrap_or_default();
333 assert_eq!(release.version, "2.1.0-beta.1");
334 Ok(())
335 }
336
337 #[test]
338 fn resolve_release_falls_back_to_top_level_version() -> Result<(), serde_json::Error> {
339 let envelope: ReleaseEnvelope = serde_json::from_value(serde_json::json!({
341 "version": "3.0.0",
342 "packageName": "pi-renamed",
343 "note": " major release "
344 }))?;
345 let stable = resolve_release(envelope.clone(), ReleaseChannel::Stable);
346 let beta = resolve_release(envelope, ReleaseChannel::Beta);
347 assert_eq!(stable.as_ref().map(|r| r.version.as_str()), Some("3.0.0"));
348 assert_eq!(beta.as_ref().map(|r| r.version.as_str()), Some("3.0.0"));
349 assert_eq!(
350 stable.and_then(|r| r.package_name),
351 Some("pi-renamed".to_owned())
352 );
353 assert_eq!(beta.and_then(|r| r.note), Some("major release".to_owned()));
354 Ok(())
355 }
356
357 #[test]
358 fn resolve_release_returns_none_for_empty_or_missing_version() -> Result<(), serde_json::Error>
359 {
360 let envelope: ReleaseEnvelope = serde_json::from_value(serde_json::json!({
362 "stable": {"version": " "}
363 }))?;
364 assert!(resolve_release(envelope, ReleaseChannel::Stable).is_none());
365
366 let envelope: ReleaseEnvelope = serde_json::from_value(serde_json::json!({}))?;
368 assert!(resolve_release(envelope.clone(), ReleaseChannel::Stable).is_none());
369 assert!(resolve_release(envelope, ReleaseChannel::Beta).is_none());
370 Ok(())
371 }
372
373 #[test]
374 fn resolve_release_trims_and_filters_empty_optional_fields() -> Result<(), serde_json::Error> {
375 let envelope: ReleaseEnvelope = serde_json::from_value(serde_json::json!({
376 "version": "1.0.0",
377 "packageName": " ",
378 "note": " "
379 }))?;
380 let result = resolve_release(envelope, ReleaseChannel::Stable);
381 assert!(result.is_some());
382 let release = result.unwrap_or_default();
383 assert_eq!(release.version, "1.0.0");
384 assert!(release.package_name.is_none());
386 assert!(release.note.is_none());
387 Ok(())
388 }
389
390 #[test]
391 fn stable_and_beta_envelopes_resolve_deterministically() -> Result<(), serde_json::Error> {
392 let envelope: ReleaseEnvelope = serde_json::from_value(serde_json::json!({
393 "stable": {"version":"2.0.0"},
394 "beta": {"version":"2.1.0-beta.1", "note":" preview "}
395 }))?;
396 let stable = resolve_release(envelope, ReleaseChannel::Stable);
397 assert_eq!(stable.map(|value| value.version), Some("2.0.0".to_owned()));
398
399 let envelope: ReleaseEnvelope = serde_json::from_value(serde_json::json!({
400 "stable": {"version":"2.0.0"},
401 "beta": {"version":"2.1.0-beta.1", "note":" preview "}
402 }))?;
403 let beta = resolve_release(envelope, ReleaseChannel::Beta);
404 assert_eq!(
405 beta.as_ref().map(|value| value.version.as_str()),
406 Some("2.1.0-beta.1")
407 );
408 assert_eq!(
409 beta.and_then(|value| value.note),
410 Some("preview".to_owned())
411 );
412 Ok(())
413 }
414
415 fn fake_endpoint(
416 body: &'static str,
417 delay: Duration,
418 ) -> Result<String, Box<dyn std::error::Error>> {
419 let listener = TcpListener::bind("127.0.0.1:0")?;
420 let address = listener.local_addr()?;
421 thread::spawn(move || {
422 let Ok((mut stream, _)) = listener.accept() else {
423 return;
424 };
425 let mut request = [0_u8; 2048];
426 let _ = stream.read(&mut request);
427 thread::sleep(delay);
428 let response = format!(
429 "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
430 body.len()
431 );
432 let _ = stream.write_all(response.as_bytes());
433 });
434 Ok(format!("http://{address}/latest"))
435 }
436
437 #[tokio::test]
438 async fn fake_http_covers_newer_equal_older_malformed_and_timeout()
439 -> Result<(), Box<dyn std::error::Error>> {
440 let client = Client::new();
441 for (body, expected) in [
442 (r#"{"version":"2.0.0"}"#, true),
443 (r#"{"version":"1.0.0"}"#, false),
444 (r#"{"version":"0.9.0"}"#, false),
445 (r#"{"wrong":"shape"}"#, false),
446 ] {
447 let endpoint = fake_endpoint(body, Duration::ZERO)?;
448 let release = check_for_new_pi_version_from(
449 &client,
450 &endpoint,
451 "1.0.0",
452 Duration::from_secs(1),
453 ReleaseChannel::Stable,
454 false,
455 )
456 .await;
457 assert_eq!(release.is_some(), expected);
458 }
459 let endpoint = fake_endpoint(r#"{"version":"2.0.0"}"#, Duration::from_millis(100))?;
460 let timed_out = check_for_new_pi_version_from(
461 &client,
462 &endpoint,
463 "1.0.0",
464 Duration::from_millis(5),
465 ReleaseChannel::Stable,
466 false,
467 )
468 .await;
469 assert!(timed_out.is_none());
470 Ok(())
471 }
472
473 #[tokio::test]
474 async fn fake_http_resolves_beta_channel_and_package_rename()
475 -> Result<(), Box<dyn std::error::Error>> {
476 let client = Client::new();
477 let endpoint = fake_endpoint(
479 r#"{"stable":{"version":"2.0.0"},"beta":{"version":"2.1.0-beta.1","packageName":"pi-new"}}"#,
480 Duration::ZERO,
481 )?;
482 let release = check_for_new_pi_version_from(
483 &client,
484 &endpoint,
485 "2.0.0",
486 Duration::from_secs(1),
487 ReleaseChannel::Beta,
488 false,
489 )
490 .await;
491 assert!(release.is_some(), "beta must be newer than stable 2.0.0");
492 let release = release.unwrap_or_default();
493 assert_eq!(release.version, "2.1.0-beta.1");
494 assert_eq!(release.package_name, Some("pi-new".to_owned()));
495 Ok(())
496 }
497
498 #[tokio::test]
499 async fn offline_injected_check_never_touches_endpoint() {
500 let result = check_for_new_pi_version_from(
501 &Client::new(),
502 "http://127.0.0.1:1/should-not-connect",
503 "1.0.0",
504 Duration::from_millis(1),
505 ReleaseChannel::Stable,
506 true,
507 )
508 .await;
509 assert!(result.is_none());
510 }
511}