use std::fs;
use std::path::{Path, PathBuf};
use toml_edit::DocumentMut;
use xshell::Shell;
use crate::environment::{get_workspace_root, CmdExt, WorkspaceManifest};
use crate::semantic_version::Version;
const COMPONENTS: &str = "rust-src,clippy,rustfmt";
const TARGET: &str = "thumbv7m-none-eabi";
const MIN_RUSTUP_VERSION_NO_UPDATE: Version = Version { major: 1, minor: 29, patch: 0 };
#[derive(Debug)]
enum ToolchainsLocation {
Workspace,
Package,
}
struct ToolchainsConfigData {
nightly: Option<String>,
stable: Option<String>,
location: ToolchainsLocation,
}
#[derive(serde::Deserialize, Default)]
struct RbmtTable {
#[serde(default)]
toolchains: Option<ToolchainsConfig>,
}
#[derive(serde::Deserialize)]
struct ToolchainsConfig {
nightly: Option<String>,
stable: Option<String>,
}
const RUSTUP_TOOLCHAIN: &str = "RUSTUP_TOOLCHAIN";
#[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)]
pub enum Toolchain {
Nightly,
Stable,
Msrv,
}
impl Toolchain {
pub fn try_read_version(self, sh: &Shell) -> Option<std::string::String> {
let config = match Self::read_toolchains_config(sh) {
Ok(c) => Some(c),
Err(e) => {
rbmt_eprintln!("Warning: Could not read toolchains config: {}", e);
None
}
};
match self {
Self::Nightly => config.and_then(|c| c.nightly),
Self::Stable => config.and_then(|c| c.stable),
Self::Msrv => match get_workspace_msrv(sh) {
Ok(msrv) => Some(msrv),
Err(e) => {
rbmt_eprintln!("Unable to determine MSRV: {}", e);
None
}
},
}
}
pub fn update_version(self, sh: &Shell) -> Result<(), Box<dyn std::error::Error>> {
let root = get_workspace_root(sh)?;
let path = root.join("Cargo.toml");
let contents = std::fs::read_to_string(&path)?;
let mut doc: toml_edit::DocumentMut = contents.parse()?;
let table = match Self::read_toolchains_config(sh)?.location {
ToolchainsLocation::Workspace =>
&mut doc["workspace"]["metadata"]["rbmt"]["toolchains"],
ToolchainsLocation::Package => &mut doc["package"]["metadata"]["rbmt"]["toolchains"],
};
let version = match self {
Self::Nightly => {
let v = Self::fetch_latest_nightly()?;
table["nightly"] = toml_edit::value(&v);
v
}
Self::Stable => {
let v = Self::fetch_latest_stable()?;
table["stable"] = toml_edit::value(&v);
v
}
Self::Msrv => return Err("Cannot update MSRV version".into()),
};
std::fs::write(&path, doc.to_string())?;
rbmt_eprintln!("Updated {:?}: {}", self, version);
Ok(())
}
fn read_toolchains_config(
sh: &Shell,
) -> Result<ToolchainsConfigData, Box<dyn std::error::Error>> {
let root = get_workspace_root(sh)?;
let contents = std::fs::read_to_string(root.join("Cargo.toml"))?;
let cargo_toml = toml::from_str::<WorkspaceManifest<RbmtTable>>(&contents)?;
if let Some(toolchains) = cargo_toml.workspace.metadata.rbmt.toolchains {
return Ok(ToolchainsConfigData {
nightly: toolchains.nightly,
stable: toolchains.stable,
location: ToolchainsLocation::Workspace,
});
}
if let Some(toolchains) = cargo_toml.package.metadata.rbmt.toolchains {
return Ok(ToolchainsConfigData {
nightly: toolchains.nightly,
stable: toolchains.stable,
location: ToolchainsLocation::Package,
});
}
Err("No [workspace.metadata.rbmt.toolchains] or [package.metadata.rbmt.toolchains] exists."
.into())
}
fn fetch_latest_nightly() -> Result<String, Box<dyn std::error::Error>> {
let manifest =
bitreq::get("https://static.rust-lang.org/dist/channel-rust-nightly.toml").send()?;
let text = manifest.as_str()?;
let parsed: toml::Value = toml::from_str(text)?;
let date = parsed
.get("date")
.and_then(|v| v.as_str())
.ok_or("Could not find 'date' field in nightly channel manifest")?;
Ok(format!("nightly-{}", date))
}
fn fetch_latest_stable() -> Result<String, Box<dyn std::error::Error>> {
let manifest =
bitreq::get("https://static.rust-lang.org/dist/channel-rust-stable.toml").send()?;
let text = manifest.as_str()?;
let parsed: toml::Value = toml::from_str(text)?;
let rustc_section = parsed
.get("pkg")
.and_then(|pkg| pkg.get("rustc"))
.ok_or("Could not find pkg.rustc section in stable channel manifest")?;
let version_str = rustc_section
.get("version")
.and_then(|v| v.as_str())
.ok_or("Could not find version field in rustc package")?;
let version = version_str
.split_whitespace()
.next()
.ok_or("Could not parse version from rustc package")?;
Ok(version.to_string())
}
}
fn rustup_version(sh: &Shell) -> Result<Version, Box<dyn std::error::Error>> {
let version_output = rbmt_cmd!(sh, "rustup --version").read()?;
if let Some(version_str) = version_output.split_whitespace().nth(1) {
if let Some(version) = Version::parse(version_str) {
return Ok(version);
}
}
Err("Could not parse rustup version".into())
}
pub fn install_toolchain(
sh: &Shell,
toolchain: &str,
force: bool,
) -> Result<(), Box<dyn std::error::Error>> {
rbmt_eprintln!("Installing toolchain {}", toolchain);
let mut install_cmd = rbmt_cmd!(
sh,
"rustup toolchain install {toolchain} --component {COMPONENTS} --target {TARGET}"
)
.arg("--no-self-update")
.env("RUSTUP_PERMIT_COPY_RENAME", "true");
if force {
install_cmd = install_cmd.arg("--force");
} else if rustup_version(sh)? >= MIN_RUSTUP_VERSION_NO_UPDATE {
install_cmd = install_cmd.arg("--no-update");
}
install_cmd.run_with_capture()?;
Ok(())
}
pub fn prepare_toolchain(
sh: &Shell,
required: Toolchain,
) -> Result<(), Box<dyn std::error::Error>> {
prepare_toolchain_with_override(sh, required, None)
}
pub fn prepare_toolchain_with_override(
sh: &Shell,
required: Toolchain,
msrv_override: Option<&str>,
) -> Result<(), Box<dyn std::error::Error>> {
if let Some(version) = &msrv_override
.filter(|_| matches!(required, Toolchain::Msrv))
.map(std::string::ToString::to_string)
.or_else(|| required.try_read_version(sh))
{
if rbmt_cmd!(sh, "rustup --version").ignore_stderr().read().is_ok() {
install_toolchain(sh, version, false)?;
sh.set_var(RUSTUP_TOOLCHAIN, version.clone());
}
}
let active_toolchain = rbmt_cmd!(sh, "rustc --version").read()?;
match required {
Toolchain::Nightly =>
if !active_toolchain.contains("nightly") {
return Err(format!("Need a nightly compiler; have {}", active_toolchain).into());
},
Toolchain::Stable =>
if active_toolchain.contains("nightly") || active_toolchain.contains("beta") {
return Err(format!("Need a stable compiler; have {}", active_toolchain).into());
},
Toolchain::Msrv => {
let active_version =
extract_version(&active_toolchain).ok_or("Could not parse rustc version")?;
let msrv_version = if let Some(override_version) = msrv_override {
rbmt_eprintln!("Using MSRV override: {}", override_version);
override_version.to_string()
} else {
let manifest_path = sh.current_dir().join("Cargo.toml");
if !manifest_path.exists() {
return Err("Not in a crate directory (no Cargo.toml found)".into());
}
get_msrv_from_manifest(sh, &manifest_path)?
};
if active_version != msrv_version {
return Err(
format!("Need Rust {} but have {}", msrv_version, active_version).into()
);
}
}
}
rbmt_eprintln!("The current toolchain is: {}", active_toolchain);
Ok(())
}
pub fn get_workspace_msrv(sh: &Shell) -> Result<String, Box<dyn std::error::Error>> {
let mut msrvs: Vec<String> =
collect_msrvs(sh)?.into_iter().filter_map(|(_, rust_version)| rust_version).collect();
msrvs.sort();
msrvs.dedup();
match msrvs.as_slice() {
[] => Err("No MSRV (rust-version) found in any Cargo.toml in the workspace".into()),
[msrv] => Ok(msrv.clone()),
_ => Err(format!("Workspace packages have conflicting MSRVs: {}", msrvs.join(", ")).into()),
}
}
fn get_msrv_from_manifest(
sh: &Shell,
manifest_path: &Path,
) -> Result<String, Box<dyn std::error::Error>> {
collect_msrvs(sh)?
.into_iter()
.find(|(path, _)| path == manifest_path)
.and_then(|(_, rust_version)| rust_version)
.ok_or_else(|| {
format!("No MSRV (rust-version) specified in {}", manifest_path.display()).into()
})
}
type ManifestMsrv = (PathBuf, Option<String>);
fn parse_rust_version_from_toml(
manifest_path: &Path,
) -> Result<Option<String>, Box<dyn std::error::Error>> {
let content = fs::read_to_string(manifest_path)?;
let doc = content.parse::<DocumentMut>()?;
if let Some(package) = doc.get("package").and_then(|p| p.as_table()) {
if let Some(rust_version) = package.get("rust-version") {
if let Some(version_str) = rust_version.as_str() {
return Ok(Some(version_str.to_string()));
}
}
}
Ok(None)
}
fn collect_msrvs(sh: &Shell) -> Result<Vec<ManifestMsrv>, Box<dyn std::error::Error>> {
let metadata = rbmt_cmd!(sh, "cargo metadata --format-version 1 --no-deps").read()?;
let data: serde_json::Value = serde_json::from_str(&metadata)?;
Ok(data["packages"]
.as_array()
.map(|packages| {
packages
.iter()
.filter_map(|pkg| {
let manifest_path = PathBuf::from(pkg["manifest_path"].as_str()?);
let rust_version = pkg["rust_version"]
.as_str()
.map(str::to_string)
.or_else(|| parse_rust_version_from_toml(&manifest_path).ok().flatten());
Some((manifest_path, rust_version))
})
.collect()
})
.unwrap_or_default())
}
fn extract_version(rustc_version: &str) -> Option<&str> {
rustc_version.split_whitespace().find_map(|part| {
let version_part = part.split('-').next()?;
let parts: Vec<&str> = version_part.split('.').collect();
if parts.len() == 3 && parts.iter().all(|p| p.chars().all(|c| c.is_ascii_digit())) {
Some(version_part)
} else {
None
}
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_extract_version() {
assert_eq!(extract_version("rustc 1.74.0 (79e9716c9 2023-11-13)"), Some("1.74.0"));
assert_eq!(extract_version("rustc 1.75.0-nightly (12345abcd 2023-11-20)"), Some("1.75.0"));
assert_eq!(extract_version("rustc 1.74.0"), Some("1.74.0"));
assert_eq!(extract_version("rustc unknown version"), None);
assert_eq!(extract_version("no version here"), None);
}
}