From b4ebf633358877dbbcee42d6c34672d42d25f956 Mon Sep 17 00:00:00 2001 From: Patrick Meinecke Date: Tue, 19 Mar 2019 16:36:27 -0400 Subject: [PATCH] Improve type inference of array literals and foreach statement variables (#8100) Improve type inference for foreach statement variables by: Inferring strongly typed arrays from explicit array and array literal expressions when elements are of the same inferred type Fix detection of foreach variable declaration. The previous logic was to check if the variable expression's start offset was after the end offset of the foreach statement, which will never be true in the body Improve inference of what type the "Condition" of a foreach statement will enumerate as --- .../engine/parser/TypeInferenceVisitor.cs | 234 +++++++++++++++++- .../engine/Api/TypeInference.Tests.ps1 | 135 +++++++++- 2 files changed, 361 insertions(+), 8 deletions(-) diff --git a/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs b/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs index 31ce46f385..a3c095b493 100644 --- a/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs +++ b/src/System.Management.Automation/engine/parser/TypeInferenceVisitor.cs @@ -514,12 +514,23 @@ namespace System.Management.Automation object ICustomAstVisitor.VisitArrayExpression(ArrayExpressionAst arrayExpressionAst) { - return new[] { new PSTypeName(typeof(object[])) }; + if (arrayExpressionAst.SubExpression.Statements.Count == 0) + { + return new[] { new PSTypeName(typeof(object[])) }; + } + + return new[] { GetArrayType(InferTypes(arrayExpressionAst.SubExpression)) }; } object ICustomAstVisitor.VisitArrayLiteral(ArrayLiteralAst arrayLiteralAst) { - return new[] { new PSTypeName(typeof(object[])) }; + var inferredElementTypes = new List(); + foreach (ExpressionAst expression in arrayLiteralAst.Elements) + { + inferredElementTypes.AddRange(InferTypes(expression)); + } + + return new[] { GetArrayType(inferredElementTypes) }; } object ICustomAstVisitor.VisitHashtable(HashtableAst hashtableAst) @@ -1971,9 +1982,22 @@ namespace System.Management.Automation int startOffset = variableExpressionAst.Extent.StartOffset; var targetAsts = (List)AstSearcher.FindAll( parent, - ast => (ast is ParameterAst || ast is AssignmentStatementAst || ast is ForEachStatementAst || ast is CommandAst) - && variableExpressionAst.AstAssignsToSameVariable(ast) - && ast.Extent.EndOffset < startOffset, + ast => + { + if (ast is ParameterAst || ast is AssignmentStatementAst || ast is CommandAst) + { + return variableExpressionAst.AstAssignsToSameVariable(ast) + && ast.Extent.EndOffset < startOffset; + } + + if (ast is ForEachStatementAst) + { + return variableExpressionAst.AstAssignsToSameVariable(ast) + && ast.Extent.StartOffset < startOffset; + } + + return false; + }, searchNestedScriptBlocks: true); foreach (var ast in targetAsts) @@ -2006,7 +2030,8 @@ namespace System.Management.Automation var foreachAst = targetAsts.OfType().FirstOrDefault(); if (foreachAst != null) { - inferredTypes.AddRange(InferTypes(foreachAst.Condition)); + inferredTypes.AddRange( + GetInferredEnumeratedTypes(InferTypes(foreachAst.Condition))); return; } @@ -2040,6 +2065,177 @@ namespace System.Management.Automation } } + /// + /// Gets the most specific array type possible from a group of inferred types. + /// + /// The inferred types all the items in the array. + /// The inferred strongly typed array type. + private PSTypeName GetArrayType(IEnumerable inferredTypes) + { + PSTypeName foundType = null; + foreach (PSTypeName inferredType in inferredTypes) + { + if (inferredType.Type == null) + { + return new PSTypeName(typeof(object[])); + } + + // IEnumerable<>.GetEnumerator and IDictionary.GetEnumerator will always be + // inferred as multiple types due to explicit implementations, so if we find + // one then assume the rest are also enumerators. + if (typeof(IEnumerator).IsAssignableFrom(inferredType.Type)) + { + foundType = inferredType; + break; + } + + if (foundType == null) + { + foundType = inferredType; + continue; + } + + // If there are mixed types then fall back to object[]. + if (foundType.Type != inferredType.Type) + { + return new PSTypeName(typeof(object[])); + } + } + + if (foundType == null) + { + return new PSTypeName(typeof(object[])); + } + + if (foundType.Type.IsArray) + { + return foundType; + } + + Type enumeratedItemType = GetMostSpecificEnumeratedItemType(foundType.Type); + if (enumeratedItemType != null) + { + return new PSTypeName(enumeratedItemType.MakeArrayType()); + } + + return new PSTypeName(foundType.Type.MakeArrayType()); + } + + /// + /// Gets the most specific type item type from a type that is potentially enumerable. + /// + /// The type to infer enumerated item type from. + /// The inferred enumerated item type. + private Type GetMostSpecificEnumeratedItemType(Type enumerableType) + { + if (enumerableType.IsArray) + { + return enumerableType.GetElementType(); + } + + // These types implement IEnumerable, but we intentionally do not enumerate them. + if (enumerableType == typeof(string) || + typeof(IDictionary).IsAssignableFrom(enumerableType) || + typeof(Xml.XmlNode).IsAssignableFrom(enumerableType)) + { + return enumerableType; + } + + if (enumerableType == typeof(Data.DataTable)) + { + return typeof(Data.DataRow); + } + + bool hasSeenNonGeneric = false; + bool hasSeenDictionaryEnumerator = false; + Type collectionInterface = GetGenericCollectionLikeInterface( + enumerableType, + ref hasSeenNonGeneric, + ref hasSeenDictionaryEnumerator); + + if (collectionInterface != null) + { + return collectionInterface.GetGenericArguments()[0]; + } + + foreach (Type interfaceType in enumerableType.GetInterfaces()) + { + collectionInterface = GetGenericCollectionLikeInterface( + interfaceType, + ref hasSeenNonGeneric, + ref hasSeenDictionaryEnumerator); + + if (collectionInterface != null) + { + return collectionInterface.GetGenericArguments()[0]; + } + } + + if (hasSeenDictionaryEnumerator) + { + return typeof(DictionaryEntry); + } + + if (hasSeenNonGeneric) + { + return typeof(object); + } + + return null; + } + + /// + /// Determines if the interface can be used to infer a specific enumerated type. + /// + /// The interface to test. + /// + /// A reference to a value indicating whether a non-generic enumerable type has been + /// seen. If is a non-generic enumerable type this + /// value will be set to . + /// + /// + /// A reference to a value indicating whether has been + /// seen. If is a this + /// value will be set to . + /// + /// + /// The value of if it can be used to infer a specific + /// enumerated type, otherwise . + /// + private Type GetGenericCollectionLikeInterface( + Type interfaceType, + ref bool hasSeenNonGeneric, + ref bool hasSeenDictionaryEnumerator) + { + if (!interfaceType.IsInterface) + { + return null; + } + + if (interfaceType.IsConstructedGenericType) + { + Type openGeneric = interfaceType.GetGenericTypeDefinition(); + if (openGeneric == typeof(IEnumerator<>) || + openGeneric == typeof(IEnumerable<>)) + { + return interfaceType; + } + } + + if (interfaceType == typeof(IDictionaryEnumerator)) + { + hasSeenDictionaryEnumerator = true; + } + + if (interfaceType == typeof(IEnumerator) || + interfaceType == typeof(IEnumerable)) + { + hasSeenNonGeneric = true; + } + + return null; + } + private IEnumerable InferTypeFrom(IndexExpressionAst indexExpressionAst) { var targetTypes = InferTypes(indexExpressionAst.Target); @@ -2102,6 +2298,32 @@ namespace System.Management.Automation } } + /// + /// Infers the types as if they were enumerated. For example, a + /// of type would be returned as . + /// + /// + /// The potentially enumerable types to infer enumerated type from. + /// + /// The enumerated item types. + private IEnumerable GetInferredEnumeratedTypes(IEnumerable enumerableTypes) + { + foreach (PSTypeName maybeEnumerableType in enumerableTypes) + { + Type type = maybeEnumerableType.Type; + if (type == null) + { + yield return maybeEnumerableType; + continue; + } + + Type enumeratedItemType = GetMostSpecificEnumeratedItemType(type); + yield return enumeratedItemType == null + ? maybeEnumerableType + : new PSTypeName(enumeratedItemType); + } + } + private void GetInferredTypeFromScriptBlockParameter(AstParameterArgumentPair argument, List inferredTypes) { var argumentPair = argument as AstPair; diff --git a/test/powershell/engine/Api/TypeInference.Tests.ps1 b/test/powershell/engine/Api/TypeInference.Tests.ps1 index 5cf7dae9fd..c31cff4de3 100644 --- a/test/powershell/engine/Api/TypeInference.Tests.ps1 +++ b/test/powershell/engine/Api/TypeInference.Tests.ps1 @@ -66,16 +66,102 @@ Describe "Type inference Tests" -tags "CI" { $res.Name | Should -Be 'System.object[]' } + It "Infers type from array expression with a single statement" { + $res = [AstTypeInference]::InferTypeOf( { @('test') }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.String[]' + } + + It "Infers type from array expression with multiple statements" { + $res = [AstTypeInference]::InferTypeOf( { @('test'; 'second test') }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.String[]' + } + + It "Infers type from array expression with mixed types" { + $res = [AstTypeInference]::InferTypeOf( { @('test'; 1) }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.Object[]' + } + + It "Infers type from array expression with nested arrays" { + $res = [AstTypeInference]::InferTypeOf( { @(@('test'); @('test2')) }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.String[]' + } + + It "Infers type from array expression with a non-generic dictionary enumerator" { + $res = [AstTypeInference]::InferTypeOf( { @(@{}.GetEnumerator()) }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.Collections.DictionaryEntry[]' + } + + It "Infers type from array expression with a generic dictionary enumerator" { + $res = [AstTypeInference]::InferTypeOf( { @([Dictionary[int, string]]::new().GetEnumerator()) }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be ([KeyValuePair[int, string][]].FullName) + } + + It "Infers type from array expression with nested non-array collections" { + $res = [AstTypeInference]::InferTypeOf( { + $list = [List[string]]::new() + $list2 = [List[string]]::new() + @($list; $list2) + }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.String[]' + } + It "Infers type from Array literal" { $res = [AstTypeInference]::InferTypeOf( { , 1 }.Ast) $res.Count | Should -Be 1 - $res.Name | Should -Be 'System.object[]' + $res.Name | Should -Be 'System.Int32[]' + } + + It "Infers type from Array literal with multiple elements" { + $res = [AstTypeInference]::InferTypeOf( { 0, 1 }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.Int32[]' + } + + It "Infers type from Array literal with mixed types" { + $res = [AstTypeInference]::InferTypeOf( { 'test', 1 }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.Object[]' + } + + It "Infers type from Array literal with nested arrays" { + $res = [AstTypeInference]::InferTypeOf( { @('test'), @('test2') }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.String[]' + } + + It "Infers type from array expression with a non-generic dictionary enumerator" { + $res = [AstTypeInference]::InferTypeOf( { , @{}.GetEnumerator() }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.Collections.DictionaryEntry[]' + } + + It "Infers type from array expression with a generic dictionary enumerator" { + $res = [AstTypeInference]::InferTypeOf( { , [Dictionary[int, string]]::new().GetEnumerator() }.Ast) + $res.Count | Should -Be 1 + $res.Name | Should -Be ([KeyValuePair[int, string][]].FullName) + } + + It "Infers type from Array literal with nested non-array collections" { + $res = [AstTypeInference]::InferTypeOf( { + $list = [List[string]]::new() + $list2 = [List[string]]::new() + $list, $list2 + }.Ast.EndBlock.Statements[2].PipelineElements[0].Expression) + $res.Count | Should -Be 1 + $res.Name | Should -Be 'System.String[]' } It "Infers type from array IndexExpresssion" { $res = [AstTypeInference]::InferTypeOf( { (1, 2, 3)[0] }.Ast) $res.Count | Should -Be 1 - $res.Name | Should -Be 'System.object' + $res.Name | Should -Be 'System.Int32' } It "Infers type from generic container IndexExpression" { @@ -566,6 +652,51 @@ Describe "Type inference Tests" -tags "CI" { } } + It 'Infers type of a Foreach statement current value variable' { + $res = [AstTypeInference]::InferTypeOf( { + foreach ($intValue in 1, 2, 3) { + $intValue + } + }.Ast.EndBlock.Statements[0].Body.Statements[0].PipelineElements[0].Expression) + + $res.Count | Should -Be 1 + $res.Name | Should -BeExactly 'System.Int32' + } + + It 'Infers type of a Foreach statement current value variable with hashtable enumerator' { + $res = [AstTypeInference]::InferTypeOf( { + foreach ($dictionaryEntry in @{}.GetEnumerator()) { + $dictionaryEntry + } + }.Ast.EndBlock.Statements[0].Body.Statements[0].PipelineElements[0].Expression) + + $res.Count | Should -Be 2 + $res.Name | Should -Be 'System.Collections.DictionaryEntry', 'System.Object' + } + + It 'Infers type of a Foreach statement current value variable with dictionary enumerator' { + $res = [AstTypeInference]::InferTypeOf( { + foreach ($keyValuePair in [Dictionary[int, string]]::new().GetEnumerator()) { + $keyValuePair + } + }.Ast.EndBlock.Statements[0].Body.Statements[0].PipelineElements[0].Expression) + + $res.Count | Should -Be 3 + $res.Name | Should -Be ([KeyValuePair[int, string]].FullName), 'System.Object', 'System.Collections.DictionaryEntry' + } + + It 'Infers type of a Foreach statement current value variable with generic IEnumerable' { + $res = [AstTypeInference]::InferTypeOf( { + $debugger = [Debugger]$Host.Runspace.Debugger + foreach ($subscriber in $debugger.GetCallStack()) { + $subscriber + } + }.Ast.EndBlock.Statements[1].Body.Statements[0].PipelineElements[0].Expression) + + $res.Count | Should -Be 1 + $res.Name | Should -BeExactly 'System.Management.Automation.CallStackFrame' + } + It 'Infers type from While statement' { $res = [AstTypeInference]::InferTypeOf( { while ($true) {