diff --git a/src/cli/proxy_mode.rs b/src/cli/proxy_mode.rs index 6768c54741..64204cc285 100644 --- a/src/cli/proxy_mode.rs +++ b/src/cli/proxy_mode.rs @@ -5,7 +5,7 @@ use crate::{ command::run_command_for_dir, config::{ActiveSource, Cfg}, process::Process, - toolchain::ResolvableLocalToolchainName, + toolchain::{Override, ResolvableLocalToolchainName}, }; #[tracing::instrument(level = "trace", skip(process))] @@ -25,7 +25,7 @@ pub async fn main( .as_ref() .map(|arg| arg.to_string_lossy()) .filter(|arg| arg.starts_with('+')) - .map(|name| ResolvableLocalToolchainName::from_str(&name[1..])) + .map(|name| Override::::from_str(&name[1..])) .transpose()?; // Build command args now while we know whether or not to skip arg 1. @@ -38,9 +38,10 @@ pub async fn main( let (toolchain, source) = cfg .local_toolchain(match toolchain { Some(name) => Some(( - name.resolve(&cfg.default_host_tuple()?)?, + name.resolve(&cfg)?.resolve(&cfg.default_host_tuple()?)?, ActiveSource::CommandLine, )), + None => None, }) .await?; diff --git a/src/cli/rustup_mode.rs b/src/cli/rustup_mode.rs index 06f14cd038..dcf16f9586 100644 --- a/src/cli/rustup_mode.rs +++ b/src/cli/rustup_mode.rs @@ -58,8 +58,8 @@ use crate::{ process::{ColorableTerminal, Process}, toolchain::{ CustomToolchainName, DistributableToolchain, LocalToolchainName, - MaybeResolvableToolchainName, ResolvableLocalToolchainName, ResolvableToolchainName, - Toolchain, ToolchainName, + MaybeResolvableToolchainName, Override, ResolvableLocalToolchainName, + ResolvableToolchainName, Toolchain, ToolchainName, }, utils::{self, ExitCode}, }; @@ -107,16 +107,18 @@ struct Rustup { value_parser = plus_toolchain_value_parser, value_hint = ValueHint::Other, )] - plus_toolchain: Option, + plus_toolchain: Option>, #[command(subcommand)] subcmd: Option, } -fn plus_toolchain_value_parser(s: &str) -> clap::error::Result { +fn plus_toolchain_value_parser( + s: &str, +) -> clap::error::Result> { use clap::{Error, error::ErrorKind}; if let Some(stripped) = s.strip_prefix('+') { - ResolvableToolchainName::from_str(stripped) + Override::::from_str(stripped) .map_err(|e| Error::raw(ErrorKind::InvalidValue, e)) } else { Err(Error::raw( @@ -159,7 +161,7 @@ enum RustupSubcmd { #[command(after_help = default_help())] Default { #[arg(help = maybe_resolvable_toolchain_arg_help())] - toolchain: Option, + toolchain: Option>, /// Install toolchains that require an emulator. See https://github.com/rust-lang/rustup/wiki/Non-host-toolchains #[arg(long)] @@ -613,7 +615,7 @@ enum OverrideSubcmd { #[command(alias = "add")] Set { #[arg(help = resolvable_toolchain_arg_help())] - toolchain: ResolvableToolchainName, + toolchain: Override, /// Path to the directory #[arg(long)] @@ -924,13 +926,13 @@ fn completion_command(cfg: &Cfg<'_>) -> clap::Command { async fn default_( cfg: &Cfg<'_>, - toolchain: Option, + toolchain: Option>, force_non_host: bool, ) -> anyhow::Result { common::warn_if_host_is_emulated(cfg.process); if let Some(toolchain) = toolchain { - match toolchain.to_owned() { + match toolchain.resolve(cfg)? { MaybeResolvableToolchainName::None => { cfg.set_default(None)?; } @@ -1180,7 +1182,7 @@ async fn update( )?; if opts.r#override { - cfg.make_override(&cfg.current_dir, &name.clone().into())?; + cfg.make_override(&cfg.current_dir, &name)?; } if opts.default @@ -1794,10 +1796,13 @@ fn pin_active_toolchain(qualified: bool, cfg: &Cfg<'_>) -> anyhow::Result, - toolchain: ResolvableToolchainName, + toolchain: Override, path: Option<&Path>, ) -> anyhow::Result { - let toolchain_name = toolchain.clone().resolve(&cfg.default_host_tuple()?)?; + let resolved_toolchain = toolchain.clone().resolve(cfg)?; + let toolchain_name = resolved_toolchain + .clone() + .resolve(&cfg.default_host_tuple()?)?; match Toolchain::new(cfg, toolchain_name.clone().into()) { Ok(_) => {} Err(e @ RustupError::ToolchainNotInstalled { .. }) => match &toolchain_name { @@ -1955,7 +1960,10 @@ async fn display_version(cfg: &mut Cfg<'_>) -> anyhow::Result<()> { cfg.toolchain_override = cfg .process .args() - .find_map(|arg| arg.strip_prefix('+').map(ResolvableToolchainName::from_str)) + .find_map(|arg| { + arg.strip_prefix('+') + .map(Override::::from_str) + }) .transpose()?; match cfg.maybe_ensure_active_toolchain(None).await { diff --git a/src/config.rs b/src/config.rs index 20ab15e72c..ced789fdff 100644 --- a/src/config.rs +++ b/src/config.rs @@ -25,12 +25,30 @@ use crate::{ process::Process, settings::{MetadataVersion, Settings, SettingsFile}, toolchain::{ - CustomToolchainName, DistributableToolchain, LocalToolchainName, PathBasedToolchainName, - ResolvableLocalToolchainName, ResolvableToolchainName, Toolchain, ToolchainName, + CustomToolchainName, DistributableToolchain, LocalToolchainName, Override, + PathBasedToolchainName, ResolvableLocalToolchainName, ResolvableToolchainName, Toolchain, + ToolchainAlias, ToolchainName, }, utils, }; +impl Override +where + T: From, +{ + pub(crate) fn resolve(self, cfg: &Cfg<'_>) -> anyhow::Result { + match self { + Self::Aliased(ToolchainAlias::Default) => { + let default = cfg + .get_default_resolvable()? + .ok_or_else(|| no_toolchain_error(cfg.process))?; + Ok(T::from(default)) + } + Self::Explicit(value) => Ok(value), + } + } +} + #[derive(Debug, ThisError)] enum OverrideFileConfigError { #[error( @@ -199,7 +217,7 @@ pub(crate) enum OverrideCfg { impl OverrideCfg { fn from_file(cfg: &Cfg<'_>, file: OverrideFile) -> anyhow::Result { let toolchain_name = match (file.toolchain.channel, file.toolchain.path) { - (Some(name), None) => ResolvableToolchainName::from_str(&name)?, + (Some(name), None) => Override::::from_str(&name)?, (None, Some(path)) => { if file.toolchain.targets.is_some() || file.toolchain.components.is_some() @@ -226,10 +244,12 @@ impl OverrideCfg { path.display() ) } - (None, None) => cfg - .get_default_resolvable()? - .ok_or_else(|| no_toolchain_error(cfg.process))?, + (None, None) => Override::Explicit( + cfg.get_default_resolvable()? + .ok_or_else(|| no_toolchain_error(cfg.process))?, + ), }; + let toolchain_name = toolchain_name.resolve(cfg)?; Ok(match toolchain_name { ResolvableToolchainName::Official(desc) => Self::Official { toolchain: desc, @@ -321,8 +341,8 @@ pub(crate) struct Cfg<'a> { pub toolchains_dir: PathBuf, update_hash_dir: PathBuf, pub download_dir: PathBuf, - pub toolchain_override: Option, - env_override: Option, + pub toolchain_override: Option>, + env_override: Option>, pub(crate) dist_root_server: String, pub dist_root_url: String, pub quiet: bool, @@ -383,7 +403,7 @@ impl<'a> Cfg<'a> { // Environment override let env_override = match &process.var_opt("RUSTUP_TOOLCHAIN")? { - Some(tc) => Some(ResolvableLocalToolchainName::from_str(tc)?), + Some(tc) => Some(Override::::from_str(tc)?), None => None, }; @@ -646,7 +666,7 @@ impl<'a> Cfg<'a> { let override_config: Option<(OverrideCfg, ActiveSource)> = // First check +toolchain override from the command line if let Some(name) = &self.toolchain_override { - Some((name.clone().into(), ActiveSource::CommandLine)) + Some((name.clone().resolve(self)?.into(), ActiveSource::CommandLine)) } // Then check the RUSTUP_TOOLCHAIN environment variable else if let Some(name) = &self.env_override { @@ -654,7 +674,7 @@ impl<'a> Cfg<'a> { // custom, distributable, and absolute path toolchains otherwise // rustup's export of a RUSTUP_TOOLCHAIN when running a process will // error when a nested rustup invocation occurs - Some((name.clone().into(), ActiveSource::Environment)) + Some((name.clone().resolve(self)?.into(), ActiveSource::Environment)) } // Then walk up the directory tree from 'path' looking for either the // directory in the override database, or a `rust-toolchain{.toml}` file, @@ -668,7 +688,6 @@ impl<'a> Cfg<'a> { else { None }; - Ok(override_config) } @@ -684,7 +703,9 @@ impl<'a> Cfg<'a> { if let Some(name) = settings.dir_override(d) { let source = ActiveSource::OverrideDb(d.to_owned()); return Ok(Some(( - ResolvableToolchainName::from_str(&name)?.into(), + Override::::from_str(&name)? + .resolve(self)? + .into(), source, ))); } @@ -735,15 +756,15 @@ impl<'a> Cfg<'a> { } })?; if let Some(toolchain_name_str) = &override_file.toolchain.channel { - let toolchain_name = ResolvableToolchainName::from_str( - toolchain_name_str.as_str(), - ) - .map_err(|_| { - anyhow!( - "invalid toolchain name detected in override file '{}'", - toolchain_file.display() - ) - })?; + let toolchain_override = + Override::::from_str(toolchain_name_str.as_str()) + .map_err(|_| { + anyhow!( + "invalid toolchain name detected in override file '{}'", + toolchain_file.display() + ) + })?; + let toolchain_name = toolchain_override.resolve(self)?; let default_host = default_host_tuple(settings, self.process); // Do not permit architecture/os selection in channels as // these are host specific and toolchain files are portable. @@ -1054,7 +1075,7 @@ impl<'a> Cfg<'a> { pub(crate) fn make_override( &self, path: &Path, - toolchain: &ResolvableToolchainName, + toolchain: &impl Display, ) -> anyhow::Result<()> { self.settings_file.with_mut(|s| { s.add_override(path, toolchain.to_string()); @@ -1274,7 +1295,7 @@ pub(crate) fn default_host_tuple(s: &Settings, process: &Process) -> TargetTuple .unwrap_or_else(|| TargetTuple::from_host_or_build(process)) } -fn no_toolchain_error(process: &Process) -> anyhow::Error { +pub(crate) fn no_toolchain_error(process: &Process) -> anyhow::Error { RustupError::ToolchainNotSelected(process.name().unwrap_or_else(|| "Rust".into())).into() } diff --git a/src/toolchain.rs b/src/toolchain.rs index 96eaa0a675..e9b653f177 100644 --- a/src/toolchain.rs +++ b/src/toolchain.rs @@ -38,8 +38,8 @@ pub(crate) use distributable::DistributableToolchain; mod names; pub(crate) use names::{ CustomToolchainName, LocalToolchainName, MaybeOfficialToolchainName, - MaybeResolvableToolchainName, PathBasedToolchainName, ResolvableLocalToolchainName, - ResolvableToolchainName, ToolchainName, + MaybeResolvableToolchainName, Override, PathBasedToolchainName, ResolvableLocalToolchainName, + ResolvableToolchainName, ToolchainAlias, ToolchainName, }; /// A toolchain installed on the local disk diff --git a/src/toolchain/names.rs b/src/toolchain/names.rs index 3615427c09..f55b2a95aa 100644 --- a/src/toolchain/names.rs +++ b/src/toolchain/names.rs @@ -74,6 +74,65 @@ pub enum InvalidName { DashPrefix(String), } +/// An alias for a toolchain name. +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)] +pub enum ToolchainAlias { + ///Refers to rustup's configured default toolchain + Default, +} + +impl FromStr for ToolchainAlias { + type Err = InvalidName; + + fn from_str(value: &str) -> Result { + match value { + "default" => Ok(Self::Default), + _ => Err(InvalidName::ToolchainName(value.into())), + } + } +} + +impl Display for ToolchainAlias { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Default => write!(f, "default"), + } + } +} + +/// A wrapper for types that can be overridden by an alias. +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)] +pub enum Override { + Aliased(ToolchainAlias), + Explicit(T), +} + +impl FromStr for Override +where + T::Err: Into, +{ + type Err = InvalidName; + + fn from_str(value: &str) -> Result { + if let Ok(alias) = ToolchainAlias::from_str(value) { + return Ok(Self::Aliased(alias)); + } + match T::from_str(value) { + Ok(t) => Ok(Self::Explicit(t)), + Err(e) => Err(e.into()), + } + } +} + +impl Display for Override { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Aliased(a) => write!(f, "{a}"), + Self::Explicit(t) => write!(f, "{t}"), + } + } +} + /// A toolchain name from user input. #[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)] pub(crate) enum ResolvableToolchainName { @@ -155,6 +214,12 @@ impl Display for MaybeResolvableToolchainName { } } +impl From for MaybeResolvableToolchainName { + fn from(value: ResolvableToolchainName) -> Self { + Self::Some(value) + } +} + /// ResolvableToolchainName + none, for overriding default-has-a-value /// situations in the CLI with an official toolchain name or none #[derive(Debug, Clone)] @@ -293,6 +358,12 @@ impl Display for ResolvableLocalToolchainName { } } +impl From for ResolvableLocalToolchainName { + fn from(value: ResolvableToolchainName) -> Self { + Self::Named(value) + } +} + /// LocalToolchainName can be used in calls to Cfg that alter configuration, /// like setting overrides, or that depend on configuration, like calculating /// the toolchain directory. It is not used to model the RUSTUP_TOOLCHAIN diff --git a/tests/suite/cli_rustup.rs b/tests/suite/cli_rustup.rs index bdd6f3fa0c..5ad7f8e66e 100644 --- a/tests/suite/cli_rustup.rs +++ b/tests/suite/cli_rustup.rs @@ -359,6 +359,82 @@ rustc-[HOST_TUPLE] "#]]); } +#[tokio::test] +async fn default_alias_uses_configured_default() { + let cx = CliTestContext::new(Scenario::SimpleV2).await; + + cx.config + .expect(["rustup", "default", "beta"]) + .await + .is_ok(); + cx.config + .expect(["rustup", "default", "default"]) + .await + .is_ok(); + cx.config + .expect(["rustup", "show"]) + .await + .with_stdout(snapbox::str![[r#" +... +beta-[HOST_TUPLE] (active, default) +... +"#]]) + .is_ok(); +} + +#[tokio::test] +async fn proxy_default_alias_uses_configured_default() { + let cx = CliTestContext::new(Scenario::SimpleV2).await; + + cx.config + .expect(["rustup", "default", "beta"]) + .await + .is_ok(); + cx.config + .expect(["rustc", "+default", "--version"]) + .await + .with_stdout(snapbox::str![[r#" +1.2.0 (hash-beta-1.2.0) + +"#]]) + .is_ok(); +} + +#[tokio::test] +async fn default_alias_directory_override_follows_default() { + let cx = CliTestContext::new(Scenario::SimpleV2).await; + + cx.config + .expect(["rustup", "default", "beta"]) + .await + .is_ok(); + cx.config + .expect(["rustup", "override", "set", "default"]) + .await + .is_ok(); + cx.config + .expect(["rustc", "--version"]) + .await + .with_stdout(snapbox::str![[r#" +1.2.0 (hash-beta-1.2.0) + +"#]]) + .is_ok(); + + cx.config + .expect(["rustup", "default", "nightly"]) + .await + .is_ok(); + cx.config + .expect(["rustc", "--version"]) + .await + .with_stdout(snapbox::str![[r#" +1.3.0 (hash-nightly-2) + +"#]]) + .is_ok(); +} + #[tokio::test] async fn default_typo_guess() { let cx = CliTestContext::new(Scenario::SimpleV2).await; @@ -3601,6 +3677,24 @@ async fn env_override_beats_file_override() { .is_ok(); } +#[tokio::test] +async fn env_override_default_uses_configured_default() { + let cx = CliTestContext::new(Scenario::SimpleV2).await; + cx.config + .expect(["rustup", "default", "stable"]) + .await + .is_ok(); + + cx.config + .expect_with_env(["rustc", "--version"], [("RUSTUP_TOOLCHAIN", "default")]) + .await + .with_stdout(snapbox::str![[r#" +1.1.0 (hash-stable-1.1.0) + +"#]]) + .is_ok(); +} + #[tokio::test] async fn plus_override_beats_file_override() { let cx = CliTestContext::new(Scenario::SimpleV2).await; @@ -4145,6 +4239,26 @@ error: rustup could not choose a version of rustc to run, because one wasn't spe .is_ok(); } +#[tokio::test] +async fn rust_toolchain_toml_default_uses_configured_default() { + let cx = CliTestContext::new(Scenario::SimpleV2).await; + cx.config + .expect(["rustup", "default", "stable"]) + .await + .is_ok(); + + let toolchain_file = cx.config.current_dir().join("rust-toolchain.toml"); + raw::write_file(&toolchain_file, "[toolchain]\nchannel = \"default\"").unwrap(); + cx.config + .expect(["rustc", "--version"]) + .await + .with_stdout(snapbox::str![[r#" +1.1.0 (hash-stable-1.1.0) + +"#]]) + .is_ok(); +} + /// Ensures that `rust-toolchain.toml` files (with `.toml` extension) only allow TOML contents #[tokio::test] async fn only_toml_in_rust_toolchain_toml() {