use std::{
path::{Path, PathBuf},
process::{Output, Stdio},
time::Duration,
};
use anyhow::{anyhow, bail};
use futures_util::{
future::{BoxFuture, Shared},
FutureExt,
};
use serde::Deserialize;
use tokio::{process::Command, task::AbortHandle};
use super::{ComputerUseSession, State};
const CHECK_TIMEOUT: Duration = Duration::from_secs(30);
const VERSION_TIMEOUT: Duration = Duration::from_secs(10);
const STDERR_EXCERPT_BYTES: usize = 400;
#[derive(Deserialize)]
struct CheckPayload {
current_version: String,
latest_version: Option<String>,
update_available: bool,
error: Option<String>,
release_notes_url: Option<String>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) enum UpdateOutcome {
Available {
current: String,
latest: String,
notes: Option<String>,
},
UpToDate {
current: String,
},
Unavailable {
current: String,
reason: String,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) enum UpdateCheckStatus {
NotChecked,
Checking,
Checked(UpdateOutcome),
Failed(String),
}
type CheckTask = Shared<BoxFuture<'static, Result<UpdateOutcome, String>>>;
pub(super) struct RunningCheck {
task: CheckTask,
abort: AbortHandle,
}
impl Drop for RunningCheck {
fn drop(&mut self) {
self.abort.abort();
}
}
#[derive(Default)]
pub(super) enum UpdateCheck {
#[default]
NotChecked,
Checking(RunningCheck),
Done(Result<UpdateOutcome, String>),
}
impl ComputerUseSession {
pub(crate) fn start_update_check(&self) -> anyhow::Result<()> {
let state = self.state();
match &*state {
State::Installing(_) => {
bail!("Cua Driver installation is pending; check again after it finishes")
}
State::Off { .. }
| State::Connecting { .. }
| State::Connected { .. }
| State::Closing { .. } => {}
}
let driver = self
.driver_path()
.ok_or_else(|| anyhow!("Cua Driver was not detected; /computer setup installs it"))?;
let home = absolute_home().ok_or_else(|| {
anyhow!("an absolute home directory is required to check for updates")
})?;
let mut check = self.update_check();
if matches!(&*check, UpdateCheck::Checking(_)) {
return Ok(());
}
let handle = tokio::spawn(async move {
check_driver(&driver, &home)
.await
.map_err(|error| error.to_string())
});
let abort = handle.abort_handle();
let task = async move { handle.await.map_err(|error| error.to_string())? }
.boxed()
.shared();
*check = UpdateCheck::Checking(RunningCheck { task, abort });
Ok(())
}
pub(crate) fn poll_update_check(&self) -> bool {
let mut check = self.update_check();
let UpdateCheck::Checking(running) = &*check else {
return false;
};
let Some(result) = running.task.clone().now_or_never() else {
return false;
};
*check = UpdateCheck::Done(result);
true
}
pub(crate) fn update_check_status(&self) -> UpdateCheckStatus {
match &*self.update_check() {
UpdateCheck::NotChecked => UpdateCheckStatus::NotChecked,
UpdateCheck::Checking(_) => UpdateCheckStatus::Checking,
UpdateCheck::Done(Ok(outcome)) => UpdateCheckStatus::Checked(outcome.clone()),
UpdateCheck::Done(Err(error)) => UpdateCheckStatus::Failed(error.clone()),
}
}
pub(super) fn reset_update_check(&self) {
*self.update_check() = UpdateCheck::NotChecked;
}
fn update_check(&self) -> std::sync::MutexGuard<'_, UpdateCheck> {
self.inner
.update_check
.lock()
.unwrap_or_else(|error| error.into_inner())
}
}
pub(super) fn validate_version(version: &str) -> anyhow::Result<()> {
let shaped = version.starts_with(|c: char| c.is_ascii_digit())
&& !version.contains("..")
&& !version.ends_with(['.', '-'])
&& version
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '-');
if !shaped {
bail!("unexpected Cua Driver version {version:?}");
}
Ok(())
}
fn absolute_home() -> Option<PathBuf> {
crate::paths::home_dir().filter(|home| home.is_absolute())
}
async fn check_driver(driver: &Path, home: &Path) -> anyhow::Result<UpdateOutcome> {
let output = run_driver(driver, home, &["check-update", "--json"], CHECK_TIMEOUT).await?;
match serde_json::from_slice::<CheckPayload>(&output.stdout) {
Ok(payload) => outcome(payload),
Err(_) if !output.status.success() => Err(anyhow!(
"{}; older drivers do not support update checks, so update manually",
failure("cua-driver check-update", &output)
)),
Err(error) => bail!("cua-driver check-update returned unexpected output: {error}"),
}
}
fn outcome(payload: CheckPayload) -> anyhow::Result<UpdateOutcome> {
let CheckPayload {
current_version: current,
latest_version,
update_available,
error,
release_notes_url: notes,
} = payload;
validate_version(¤t)?;
Ok(match (error, update_available, latest_version) {
(Some(reason), _, _) => UpdateOutcome::Unavailable { current, reason },
(None, true, Some(latest)) => {
validate_version(&latest)?;
UpdateOutcome::Available {
current,
latest,
notes,
}
}
(None, true, None) => UpdateOutcome::Unavailable {
current,
reason: "the driver reported an update without a version".into(),
},
(None, false, _) => UpdateOutcome::UpToDate { current },
})
}
pub(super) async fn driver_version(driver: &Path) -> anyhow::Result<String> {
let home = absolute_home().ok_or_else(|| anyhow!("an absolute home directory is required"))?;
let output = run_driver(driver, &home, &["--version"], VERSION_TIMEOUT).await?;
if !output.status.success() {
return Err(failure("cua-driver --version", &output));
}
let output = String::from_utf8_lossy(&output.stdout);
let version = output
.split_whitespace()
.last()
.map(|token| token.trim_start_matches('v'))
.unwrap_or_default();
validate_version(version)
.map_err(|_| anyhow!("unexpected cua-driver --version output: {}", output.trim()))?;
Ok(version.to_owned())
}
async fn run_driver(
driver: &Path,
home: &Path,
args: &[&str],
timeout: Duration,
) -> anyhow::Result<Output> {
let mut command = Command::new(driver);
command.args(args);
super::setup::restrict_environment(&mut command, home);
command
.envs(super::policy::telemetry_environment())
.stdin(Stdio::null())
.kill_on_drop(true);
let invocation = format!("cua-driver {}", args.join(" "));
Ok(tokio::time::timeout(timeout, command.output())
.await
.map_err(|_| anyhow!("{invocation} exceeded the {}s limit", timeout.as_secs()))??)
}
fn failure(invocation: &str, output: &Output) -> anyhow::Error {
let stderr = String::from_utf8_lossy(&output.stderr);
let stderr = stderr.trim();
let start = stderr.floor_char_boundary(stderr.len().saturating_sub(STDERR_EXCERPT_BYTES));
anyhow!(
"{invocation} exited with {}: {}",
output.status,
&stderr[start..]
)
}
#[cfg(test)]
#[path = "update_tests.rs"]
mod tests;