diff --git a/src/System.Management.Automation/engine/ErrorPackage.cs b/src/System.Management.Automation/engine/ErrorPackage.cs index bfd6455592..a207125586 100644 --- a/src/System.Management.Automation/engine/ErrorPackage.cs +++ b/src/System.Management.Automation/engine/ErrorPackage.cs @@ -1670,6 +1670,19 @@ namespace System.Management.Automation } // class ErrorRecord + /// + /// Dummy generic class for type inference purposes on typed catch blocks. + /// + /// Anything that inherits Exception. + internal class ErrorRecord : ErrorRecord where TException : Exception + { + public new TException Exception { get; } + + public ErrorRecord(Exception exception, string errorId, ErrorCategory errorCategory, object targetObject) : base(exception, errorId, errorCategory, targetObject) + { + } + } + /// /// Implemented by exception classes which contain additional /// diff --git a/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs b/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs index 482d212e1b..5d295dcf12 100644 --- a/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs +++ b/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs @@ -1803,7 +1803,7 @@ namespace System.Management.Automation // $_ is special, see if we're used in a script block in some pipeline. while (parent != null) { - if (parent is ScriptBlockExpressionAst) + if (parent is ScriptBlockExpressionAst || parent is CatchClauseAst) { break; } @@ -1832,6 +1832,27 @@ namespace System.Management.Automation parent = parent.Parent; } + if (parent is CatchClauseAst catchBlock) + { + if (catchBlock.CatchTypes.Count > 0) + { + foreach (TypeConstraintAst catchType in catchBlock.CatchTypes) + { + Type exceptionType = catchType.TypeName.GetReflectionType(); + if (exceptionType != null && typeof(Exception).IsAssignableFrom(exceptionType)) + { + inferredTypes.Add(new PSTypeName(typeof(ErrorRecord<>).MakeGenericType(exceptionType))); + } + } + } + else + { + inferredTypes.Add(new PSTypeName(typeof(ErrorRecord))); + } + + return; + } + if (parent.Parent is CommandAst commandAst) { // We found a command, see if there is a previous command in the pipeline. diff --git a/test/powershell/engine/Api/TypeInference.Tests.ps1 b/test/powershell/engine/Api/TypeInference.Tests.ps1 index bdc1c25a31..5d6cf41676 100644 --- a/test/powershell/engine/Api/TypeInference.Tests.ps1 +++ b/test/powershell/engine/Api/TypeInference.Tests.ps1 @@ -926,6 +926,90 @@ Describe "Type inference Tests" -tags "CI" { $res.Name | Should -Be System.Int32 } + It 'Infers type of variable $_ in catch block' { + $variableAst = { try {} catch { $_ } }.Ast.Find({ param($a) $a -is [System.Management.Automation.Language.VariableExpressionAst] }, $true) + $res = [AstTypeInference]::InferTypeOf($variableAst) + + $res | Should -HaveCount 1 + $res.Name | Should -Be System.Management.Automation.ErrorRecord + } + + It 'Infers type of untyped $_.Exception in catch block' { + $memberAst = { try {} catch { $_.Exception } }.Ast.Find({ param($a) $a -is [System.Management.Automation.Language.MemberExpressionAst] }, $true) + $res = [AstTypeInference]::InferTypeOf($memberAst) + + $res | Should -HaveCount 1 + $res.Name | Should -Be System.Exception + } + + $catchClauseTypes = @( + @{ Type = 'System.ArgumentException' } + @{ Type = 'System.ArgumentNullException' } + @{ Type = 'System.ArgumentOutOfRangeException' } + @{ Type = 'System.Collections.Generic.KeyNotFoundException' } + @{ Type = 'System.DivideByZeroException' } + @{ Type = 'System.FormatException' } + @{ Type = 'System.IndexOutOfRangeException' } + @{ Type = 'System.InvalidOperationException' } + @{ Type = 'System.IO.DirectoryNotFoundException' } + @{ Type = 'System.IO.DriveNotFoundException' } + @{ Type = 'System.IO.FileNotFoundException' } + @{ Type = 'System.IO.PathTooLongException' } + @{ Type = 'System.Management.Automation.CommandNotFoundException' } + @{ Type = 'System.Management.Automation.JobFailedException' } + @{ Type = 'System.Management.Automation.RuntimeException' } + @{ Type = 'System.Management.Automation.ValidationMetadataException' } + @{ Type = 'System.NotImplementedException' } + @{ Type = 'System.NotSupportedException' } + @{ Type = 'System.ObjectDisposedException' } + @{ Type = 'System.OverflowException' } + @{ Type = 'System.PlatformNotSupportedException' } + @{ Type = 'System.RankException' } + @{ Type = 'System.TimeoutException' } + @{ Type = 'System.UriFormatException' } + ) + + It 'Infers type of $_.Exception in [] typed catch block' -TestCases $catchClauseTypes { + param($Type) + + $memberAst = [scriptblock]::Create("try {} catch [$Type] { `$_.Exception }").Ast.Find( + { param($a) $a -is [System.Management.Automation.Language.MemberExpressionAst] }, + $true + ) + $res = [AstTypeInference]::InferTypeOf($memberAst) + + $res | Should -HaveCount 1 + $res.Name | Should -Be $Type + } + + It 'Infers possible types of $_.Exception in multi-typed catch block' { + $memberAst = { try {} catch [System.ArgumentException], [System.NotImplementedException] { $_.Exception } }.Ast.Find( + { param($a) $a -is [System.Management.Automation.Language.MemberExpressionAst] }, + $true + ) + $res = [AstTypeInference]::InferTypeOf($memberAst) + + $res | Should -HaveCount 2 + $res[0].Name | Should -Be System.ArgumentException + $res[1].Name | Should -Be System.NotImplementedException + } + + It 'Infers type of $_.Exception in each successive catch block' { + $memberAst = { + try {} + catch [System.ArgumentException] { $_.Exception } + catch { $_.Exception } + }.Ast.FindAll( + { param($a) $a -is [System.Management.Automation.Language.MemberExpressionAst] }, + $true + ) + $res = foreach ($item in $memberAst) { [AstTypeInference]::InferTypeOf($item) } + + $res | Should -HaveCount 2 + $res[0].Name | Should -Be System.ArgumentException + $res[1].Name | Should -Be System.Exception + } + It 'Infers type of function member' { $res = [AstTypeInference]::InferTypeOf( { class X {