From 3d4e294262d2f2d7472ebd5d849b1b68fd6ef524 Mon Sep 17 00:00:00 2001 From: MartinGC94 <42123497+MartinGC94@users.noreply.github.com> Date: Mon, 25 Jul 2022 20:27:23 +0200 Subject: [PATCH] Improve type inference for `$_` (#17716) --- .../engine/parser/TypeInferenceVisitor.cs | 135 ++++++++++-------- .../engine/Api/TypeInference.Tests.ps1 | 85 +++++++++++ 2 files changed, 162 insertions(+), 58 deletions(-) diff --git a/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs b/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs index 3ace7fb3a0..25bd7f999a 100644 --- a/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs +++ b/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs @@ -1840,82 +1840,101 @@ namespace System.Management.Automation (SpecialVariables.IsUnderbar(astVariablePath.UserPath) || astVariablePath.UserPath.EqualsOrdinalIgnoreCase(SpecialVariables.PSItem))) { - // $_ is special, see if we're used in a script block in some pipeline. - while (parent != null) + // The automatic variable $_ is assigned a value in scriptblocks, Switch loops and Catch/Trap statements + // This loop will find whichever Ast that determines the value of $_ + // The value in scriptblocks is determined by the parents of that scriptblock, the only interesting scenarios are: + // 1: MemberInvocation like: $Collection.Where({$_}) + // 2: Command pipelines like: dir | where {$_} + // The value in a Switch loop is whichever item is in the condition part of the statement. + // The value in Catch/Trap statements is always an error record. + bool hasSeenScriptBlock = false; + while (parent is not null) { - if (parent is ScriptBlockExpressionAst || parent is CatchClauseAst) + if (parent is CatchClauseAst or TrapStatementAst) { break; } + else if (parent is SwitchStatementAst switchStatement) + { + parent = switchStatement.Condition; + break; + } + else if (parent is ErrorStatementAst switchErrorStatement && switchErrorStatement.Kind?.Kind == TokenKind.Switch) + { + if (switchErrorStatement.Conditions?.Count > 0) + { + parent = switchErrorStatement.Conditions[0]; + } + break; + } + else if (parent is ScriptBlockExpressionAst) + { + hasSeenScriptBlock = true; + } + else if (hasSeenScriptBlock) + { + if (parent is InvokeMemberExpressionAst invokeMember) + { + parent = invokeMember.Expression; + break; + } + else if (parent is CommandAst cmdAst && cmdAst.Parent is PipelineAst pipeline && pipeline.PipelineElements.Count > 1) + { + // We've found a pipeline with multiple commands, now we need to determine what command came before the command with the scriptblock: + // eg Get-Partition in this example: Get-Disk | Get-Partition | Where {$_} + var indexOfPreviousCommand = pipeline.PipelineElements.IndexOf(cmdAst) - 1; + if (indexOfPreviousCommand >= 0) + { + parent = pipeline.PipelineElements[indexOfPreviousCommand]; + break; + } + } + } parent = parent.Parent; } - if (parent != null) + if (parent is CatchClauseAst catchBlock) { - if (parent.Parent is CommandExpressionAst && parent.Parent.Parent is PipelineAst) + if (catchBlock.CatchTypes.Count > 0) { - // Script block in a hash table, could be something like: - // dir | ft @{ Expression = { $_ } } - - if (parent.Parent.Parent.Parent is HashtableAst) + foreach (TypeConstraintAst catchType in catchBlock.CatchTypes) { - parent = parent.Parent.Parent.Parent; - } - else if (parent.Parent.Parent.Parent is ArrayLiteralAst && parent.Parent.Parent.Parent.Parent is HashtableAst) - { - parent = parent.Parent.Parent.Parent.Parent; - } - } - - if (parent.Parent is CommandParameterAst) - { - 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 (typeof(Exception).IsAssignableFrom(exceptionType)) { - Type exceptionType = catchType.TypeName.GetReflectionType(); - if (exceptionType != null && typeof(Exception).IsAssignableFrom(exceptionType)) - { - inferredTypes.Add(new PSTypeName(typeof(ErrorRecord<>).MakeGenericType(exceptionType))); - } + inferredTypes.Add(new PSTypeName(typeof(ErrorRecord<>).MakeGenericType(exceptionType))); } } - else - { - inferredTypes.Add(new PSTypeName(typeof(ErrorRecord))); - } - - return; } - if (parent.Parent is CommandAst commandAst) + // Either no type constraint was specified, or all the specified catch types were unavailable but we still know it's an error record. + if (inferredTypes.Count == 0) { - // We found a command, see if there is a previous command in the pipeline. - PipelineAst pipelineAst = (PipelineAst)commandAst.Parent; - var previousCommandIndex = pipelineAst.PipelineElements.IndexOf(commandAst) - 1; - if (previousCommandIndex < 0) - { - return; - } - - AddInferredTypesForDollarUnderbar(pipelineAst.PipelineElements[0], inferredTypes); - - return; - } - - if (parent.Parent is InvokeMemberExpressionAst memberExpression) - { - AddInferredTypesForDollarUnderbar(memberExpression.Expression, inferredTypes); - - return; + inferredTypes.Add(new PSTypeName(typeof(ErrorRecord))); } } + else if (parent is TrapStatementAst trap) + { + if (trap.TrapType is not null) + { + Type exceptionType = trap.TrapType.TypeName.GetReflectionType(); + if (typeof(Exception).IsAssignableFrom(exceptionType)) + { + inferredTypes.Add(new PSTypeName(typeof(ErrorRecord<>).MakeGenericType(exceptionType))); + } + } + if (inferredTypes.Count == 0) + { + inferredTypes.Add(new PSTypeName(typeof(ErrorRecord))); + } + } + else if (parent is not null) + { + AddInferredTypesForDollarUnderbar(parent, inferredTypes); + } + + return; } // For certain variables, we always know their type, well at least we can assume we know. @@ -2072,7 +2091,7 @@ namespace System.Management.Automation continue; } - if (typeof(IEnumerable).IsAssignableFrom(result.Type)) + if (result.Type != typeof(string) && typeof(IEnumerable).IsAssignableFrom(result.Type)) { // We can't deduce much from IEnumerable, but we can if it's generic. var enumerableInterfaces = result.Type.GetInterfaces(); diff --git a/test/powershell/engine/Api/TypeInference.Tests.ps1 b/test/powershell/engine/Api/TypeInference.Tests.ps1 index fbcc784570..cd99abaf46 100644 --- a/test/powershell/engine/Api/TypeInference.Tests.ps1 +++ b/test/powershell/engine/Api/TypeInference.Tests.ps1 @@ -1079,6 +1079,42 @@ Describe "Type inference Tests" -tags "CI" { $res.Name | Should -Be System.Exception } + It 'Infers type of variable $_ in pipeline with more than one element' { + $memberAst = { Get-Date | New-Guid | Select-Object -Property {$_} }.Ast.Find({ param($a) $a -is [System.Management.Automation.Language.VariableExpressionAst] }, $true) + $res = [AstTypeInference]::InferTypeOf($memberAst) + + $res | Should -HaveCount 1 + $res.Name | Should -Be System.Guid + } + It 'Infers type of variable $_ in array of calculated properties' { + $variableAst = { New-TimeSpan | Select-Object -Property Day,@{n="min";e={$_}} }.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.TimeSpan + } + + It 'Infers type of variable $_ in switch statement' { + $variableAst = { + switch ("Hello","World") + { + 'Hello' + { + $_ + } + } }.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.String + } + + It 'Does not infer string in pipeline as char' { + $variableAst = { "Hello" | Select-Object -Property @{n="min";e={$_}} }.Ast.Find({ param($a) $a -is [System.Management.Automation.Language.VariableExpressionAst] }, $true) + $res = [AstTypeInference]::InferTypeOf($variableAst) + $res.Name | Should -Be System.String + } + $catchClauseTypes = @( @{ Type = 'System.ArgumentException' } @{ Type = 'System.ArgumentNullException' } @@ -1147,6 +1183,55 @@ Describe "Type inference Tests" -tags "CI" { $res[1].Name | Should -Be System.Exception } + It 'falls back to a generic ErrorRecord if catch exception type is invalid' { + $VariableAst = { + try {} + catch [ThisTypeDoesNotExist] { $_ } + }.Ast.Find( + { param($a) $a -is [System.Management.Automation.Language.VariableExpressionAst] }, + $true + ) + $res = [AstTypeInference]::InferTypeOf($VariableAst) + + $res.Name | Should -Be System.Management.Automation.ErrorRecord + } + + It 'Infers type of trap statement' { + $VariableAst = { + trap { $_ } + }.Ast.Find( + { param($a) $a -is [System.Management.Automation.Language.VariableExpressionAst] }, + $true + ) + $res = [AstTypeInference]::InferTypeOf($VariableAst) + + $res.Name | Should -Be System.Management.Automation.ErrorRecord + } + + It 'Infers type of exception in typed trap statement' { + $memberAst = { + trap [System.DivideByZeroException] { $_.Exception } + }.Ast.Find( + { param($a) $a -is [System.Management.Automation.Language.MemberExpressionAst] }, + $true + ) + $res = [AstTypeInference]::InferTypeOf($memberAst) + + $res.Name | Should -Be System.DivideByZeroException + } + + It 'falls back to a generic ErrorRecord if trap exception type is invalid' { + $VariableAst = { + trap [ThisTypeDoesNotExist] { $_ } + }.Ast.Find( + { param($a) $a -is [System.Management.Automation.Language.VariableExpressionAst] }, + $true + ) + $res = [AstTypeInference]::InferTypeOf($VariableAst) + + $res.Name | Should -Be System.Management.Automation.ErrorRecord + } + It 'Infers type of function member' { $res = [AstTypeInference]::InferTypeOf( { class X {