diff --git a/src/Microsoft.PowerShell.ConsoleHost/host/msh/CommandLineParameterParser.cs b/src/Microsoft.PowerShell.ConsoleHost/host/msh/CommandLineParameterParser.cs index a63c919334..6908320db5 100644 --- a/src/Microsoft.PowerShell.ConsoleHost/host/msh/CommandLineParameterParser.cs +++ b/src/Microsoft.PowerShell.ConsoleHost/host/msh/CommandLineParameterParser.cs @@ -1159,6 +1159,14 @@ namespace Microsoft.PowerShell showHelp: true); return false; } +#if !UNIX + // Only do the .ps1 extension check on Windows since shebang is not supported + if (!_file.EndsWith(".ps1", StringComparison.OrdinalIgnoreCase)) + { + SetCommandLineError(string.Format(CultureInfo.CurrentCulture, CommandLineParameterParserStrings.InvalidFileArgumentExtension, args[i])); + return false; + } +#endif i++; diff --git a/src/Microsoft.PowerShell.ConsoleHost/host/msh/ConsoleHost.cs b/src/Microsoft.PowerShell.ConsoleHost/host/msh/ConsoleHost.cs index 167186ae0c..e54e7b0263 100644 --- a/src/Microsoft.PowerShell.ConsoleHost/host/msh/ConsoleHost.cs +++ b/src/Microsoft.PowerShell.ConsoleHost/host/msh/ConsoleHost.cs @@ -1857,13 +1857,15 @@ namespace Microsoft.PowerShell Pipeline tempPipeline = exec.CreatePipeline(); Command c; +#if UNIX // if file doesn't have .ps1 extension, we read the contents and treat it as a script to support shebang with no .ps1 extension usage - if (!Path.GetExtension(filePath).Equals(".ps1", StringComparison.OrdinalIgnoreCase)) + if (!filePath.EndsWith(".ps1", StringComparison.OrdinalIgnoreCase)) { string script = File.ReadAllText(filePath); c = new Command(script, isScript: true, useLocalScope: false); } else +#endif { c = new Command(filePath, false, false); } diff --git a/test/powershell/Host/ConsoleHost.Tests.ps1 b/test/powershell/Host/ConsoleHost.Tests.ps1 index 28561da23c..7b7571ce17 100644 --- a/test/powershell/Host/ConsoleHost.Tests.ps1 +++ b/test/powershell/Host/ConsoleHost.Tests.ps1 @@ -147,21 +147,32 @@ Describe "ConsoleHost unit tests" -tags "Feature" { } It "-File should be default parameter" { - Set-Content -Path $testdrive/test -Value "'hello'" - $observed = & $powershell -NoProfile $testdrive/test + Set-Content -Path $testdrive/test.ps1 -Value "'hello'" + $observed = & $powershell -NoProfile $testdrive/test.ps1 $observed | Should -Be "hello" } - It "-File accepts scripts with and without .ps1 extension: " -TestCases @( - @{Filename="test.ps1"}, - @{Filename="test"} - ) { - param($Filename) + It "-File accepts scripts with .ps1 extension" { + $Filename = 'test.ps1' Set-Content -Path $testdrive/$Filename -Value "'hello'" $observed = & $powershell -NoProfile -File $testdrive/$Filename $observed | Should -Be "hello" } + It "-File accepts scripts without .ps1 extension to support shebang" -Skip:($IsWindows) { + $Filename = 'test.xxx' + Set-Content -Path $testdrive/$Filename -Value "'hello'" + $observed = & $powershell -NoProfile -File $testdrive/$Filename + $observed | Should -Be "hello" + } + + It "-File should fail for script without .ps1 extension" -Skip:(!$IsWindows) { + $Filename = 'test.xxx' + Set-Content -Path $testdrive/$Filename -Value "'hello'" + & $powershell -NoProfile -File $testdrive/$Filename > $null + $LASTEXITCODE | Should -Be 64 + } + It "-File should pass additional arguments to script" { Set-Content -Path $testdrive/script.ps1 -Value 'foreach($arg in $args){$arg}' $observed = & $powershell -NoProfile $testdrive/script.ps1 foo bar @@ -208,11 +219,8 @@ Describe "ConsoleHost unit tests" -tags "Feature" { $observed | Should -Be $BoolValue } - It "-File '' should return exit code from script" -TestCases @( - @{Filename = "test.ps1"}, - @{Filename = "test"} - ) { - param($Filename) + It "-File should return exit code from script" { + $Filename = 'test.ps1' Set-Content -Path $testdrive/$Filename -Value 'exit 123' & $powershell $testdrive/$Filename $LASTEXITCODE | Should -Be 123 diff --git a/test/xUnit/csharp/test_CommandLineParser.cs b/test/xUnit/csharp/test_CommandLineParser.cs index 41f501688f..cf2a128121 100644 --- a/test/xUnit/csharp/test_CommandLineParser.cs +++ b/test/xUnit/csharp/test_CommandLineParser.cs @@ -103,18 +103,27 @@ namespace PSTests.Parallel [Fact] public static void TestDefaultParameterIsFileName_Exist() { - var fileName = System.IO.Path.GetTempFileName(); + var tempFile = System.IO.Path.GetTempFileName(); + var tempPs1 = tempFile + ".ps1"; + File.Move(tempFile, tempPs1); var cpp = new CommandLineParameterParser(); - cpp.Parse(new string[] { fileName }); + cpp.Parse(new string[] { tempPs1 }); - Assert.False(cpp.AbortStartup); - Assert.False(cpp.NoExit); - Assert.False(cpp.ShowShortHelp); - Assert.False(cpp.ShowBanner); - Assert.Equal(CommandLineParameterParser.NormalizeFilePath(fileName), cpp.File); - Assert.Null(cpp.ErrorMessage); + try + { + Assert.False(cpp.AbortStartup); + Assert.False(cpp.NoExit); + Assert.False(cpp.ShowShortHelp); + Assert.False(cpp.ShowBanner); + Assert.Equal(CommandLineParameterParser.NormalizeFilePath(tempPs1), cpp.File); + Assert.Null(cpp.ErrorMessage); + } + finally + { + File.Delete(tempPs1); + } } [Theory] @@ -1212,7 +1221,16 @@ namespace PSTests.Parallel public class TestDataLastFile : IEnumerable { - private readonly string _fileName = Path.GetTempFileName(); + private static string _fileName + { + get + { + var tempFile = Path.GetTempFileName(); + var tempPs1 = tempFile + ".ps1"; + File.Move(tempFile, tempPs1); + return tempPs1; + } + } public IEnumerator GetEnumerator() { @@ -1230,21 +1248,28 @@ namespace PSTests.Parallel cpp.Parse(commandLine); - Assert.False(cpp.AbortStartup); - Assert.False(cpp.NoExit); - Assert.False(cpp.ShowShortHelp); - Assert.False(cpp.ShowBanner); - if (Platform.IsWindows) + try { - Assert.True(cpp.StaMode); - } - else - { - Assert.False(cpp.StaMode); - } + Assert.False(cpp.AbortStartup); + Assert.False(cpp.NoExit); + Assert.False(cpp.ShowShortHelp); + Assert.False(cpp.ShowBanner); + if (Platform.IsWindows) + { + Assert.True(cpp.StaMode); + } + else + { + Assert.False(cpp.StaMode); + } - Assert.Equal(CommandLineParameterParser.NormalizeFilePath(commandLine[commandLine.Length - 1]), cpp.File); - Assert.Null(cpp.ErrorMessage); + Assert.Equal(CommandLineParameterParser.NormalizeFilePath(commandLine[commandLine.Length - 1]), cpp.File); + Assert.Null(cpp.ErrorMessage); + } + finally + { + File.Delete(cpp.File); + } } } }