Improve variable type inference (#19830)

This commit is contained in:
MartinGC94
2025-03-13 10:04:02 +00:00
committed by GitHub
parent 0fb7e3b18c
commit becdd61b3c
3 changed files with 745 additions and 197 deletions
@@ -5551,7 +5551,7 @@ namespace System.Management.Automation
}
}
private static readonly HashSet<string> s_varModificationCommands = new(StringComparer.OrdinalIgnoreCase)
internal static readonly HashSet<string> s_varModificationCommands = new(StringComparer.OrdinalIgnoreCase)
{
"New-Variable",
"nv",
@@ -5560,13 +5560,13 @@ namespace System.Management.Automation
"sv"
};
private static readonly string[] s_varModificationParameters = new string[]
internal static readonly string[] s_varModificationParameters = new string[]
{
"Name",
"Value"
};
private static readonly string[] s_outVarParameters = new string[]
internal static readonly string[] s_outVarParameters = new string[]
{
"ErrorVariable",
"ev",
@@ -5579,13 +5579,13 @@ namespace System.Management.Automation
};
private static readonly string[] s_pipelineVariableParameters = new string[]
internal static readonly string[] s_pipelineVariableParameters = new string[]
{
"PipelineVariable",
"pv"
};
private static readonly HashSet<string> s_localScopeCommandNames = new(StringComparer.OrdinalIgnoreCase)
internal static readonly HashSet<string> s_localScopeCommandNames = new(StringComparer.OrdinalIgnoreCase)
{
"Microsoft.PowerShell.Core\\ForEach-Object",
"ForEach-Object",
@@ -1272,7 +1272,7 @@ namespace System.Management.Automation
return TypeInferenceContext.EmptyPSTypeNameArray;
}
private void InferTypesFrom(CommandAst commandAst, List<PSTypeName> inferredTypes)
private void InferTypesFrom(CommandAst commandAst, List<PSTypeName> inferredTypes, bool forRedirection = false)
{
if (commandAst.Redirections.Count > 0)
{
@@ -1282,7 +1282,7 @@ namespace System.Management.Automation
{
if (streamRedirection is FileRedirectionAst fileRedirection)
{
if (fileRedirection.FromStream is RedirectionStream.All or RedirectionStream.Output)
if (!forRedirection && fileRedirection.FromStream is RedirectionStream.All or RedirectionStream.Output)
{
// command output is redirected so it returns nothing.
return;
@@ -1952,6 +1952,36 @@ namespace System.Management.Automation
return res;
}
private static IEnumerable<PSTypeName> InferTypeFromRef(InvokeMemberExpressionAst invokeMember, ExpressionAst refArgument)
{
Type expressionClrType = (invokeMember.Expression as TypeExpressionAst)?.TypeName.GetReflectionType();
string memberName = (invokeMember.Member as StringConstantExpressionAst)?.Value;
int argumentIndex = invokeMember.Arguments.IndexOf(refArgument);
if (expressionClrType is null || string.IsNullOrEmpty(memberName) || argumentIndex == -1)
{
yield break;
}
foreach (MemberInfo memberInfo in expressionClrType.GetMember(memberName))
{
if (memberInfo.MemberType == MemberTypes.Method)
{
var methodInfo = memberInfo as MethodInfo;
ParameterInfo[] methodParams = methodInfo.GetParameters();
if (methodParams.Length < argumentIndex)
{
continue;
}
ParameterInfo paramCandidate = methodParams[argumentIndex];
if (paramCandidate.IsOut)
{
yield return new PSTypeName(paramCandidate.ParameterType.GetElementType());
}
}
}
}
private void GetTypesOfMembers(
PSTypeName thisType,
string memberName,
@@ -2248,7 +2278,7 @@ namespace System.Management.Automation
return;
}
Ast parent = variableExpressionAst.Parent;
Ast currentAst = variableExpressionAst.Parent;
if (astVariablePath.IsUnqualified &&
(SpecialVariables.IsUnderbar(astVariablePath.UserPath)
|| astVariablePath.UserPath.EqualsOrdinalIgnoreCase(SpecialVariables.PSItem)))
@@ -2261,65 +2291,65 @@ namespace System.Management.Automation
// 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)
while (currentAst is not null)
{
if (parent is CatchClauseAst or TrapStatementAst)
if (currentAst is CatchClauseAst or TrapStatementAst)
{
break;
}
else if (parent is SwitchStatementAst switchStatement
else if (currentAst is SwitchStatementAst switchStatement
&& switchStatement.Condition.Extent.EndOffset < variableExpressionAst.Extent.StartOffset)
{
parent = switchStatement.Condition;
currentAst = switchStatement.Condition;
break;
}
else if (parent is ErrorStatementAst switchErrorStatement && switchErrorStatement.Kind?.Kind == TokenKind.Switch)
else if (currentAst is ErrorStatementAst switchErrorStatement && switchErrorStatement.Kind?.Kind == TokenKind.Switch)
{
if (switchErrorStatement.Conditions?.Count > 0)
{
if (switchErrorStatement.Conditions[0].Extent.EndOffset < variableExpressionAst.Extent.StartOffset)
{
parent = switchErrorStatement.Conditions[0];
currentAst = switchErrorStatement.Conditions[0];
break;
}
else
{
// $_ is inside the condition that is being declared, eg: Get-Process | Sort-Object -Property {switch ($_.Proc<Tab>
parent = switchErrorStatement.Parent;
currentAst = switchErrorStatement.Parent;
continue;
}
}
break;
}
else if (parent is ScriptBlockExpressionAst)
else if (currentAst is ScriptBlockExpressionAst)
{
hasSeenScriptBlock = true;
}
else if (hasSeenScriptBlock)
{
if (parent is InvokeMemberExpressionAst invokeMember)
if (currentAst is InvokeMemberExpressionAst invokeMember)
{
parent = invokeMember.Expression;
currentAst = invokeMember.Expression;
break;
}
else if (parent is CommandAst cmdAst && cmdAst.Parent is PipelineAst pipeline && pipeline.PipelineElements.Count > 1)
else if (currentAst 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];
currentAst = pipeline.PipelineElements[indexOfPreviousCommand];
break;
}
}
}
parent = parent.Parent;
currentAst = currentAst.Parent;
}
if (parent is CatchClauseAst catchBlock)
if (currentAst is CatchClauseAst catchBlock)
{
if (catchBlock.CatchTypes.Count > 0)
{
@@ -2339,7 +2369,7 @@ namespace System.Management.Automation
inferredTypes.Add(new PSTypeName(typeof(ErrorRecord)));
}
}
else if (parent is TrapStatementAst trap)
else if (currentAst is TrapStatementAst trap)
{
if (trap.TrapType is not null)
{
@@ -2354,156 +2384,126 @@ namespace System.Management.Automation
inferredTypes.Add(new PSTypeName(typeof(ErrorRecord)));
}
}
else if (parent is not null)
else if (currentAst is not null)
{
inferredTypes.AddRange(GetInferredEnumeratedTypes(InferTypes(parent)));
inferredTypes.AddRange(GetInferredEnumeratedTypes(InferTypes(currentAst)));
}
return;
}
// For certain variables, we always know their type, well at least we can assume we know.
if (astVariablePath.IsUnqualified)
// Process the well known variable $this
if (astVariablePath.IsUnqualified
&& astVariablePath.UnqualifiedPath.EqualsOrdinalIgnoreCase(SpecialVariables.This)
&& (_context.CurrentTypeDefinitionAst is not null || _context.CurrentThisType is not null))
{
var isThis = astVariablePath.UserPath.EqualsOrdinalIgnoreCase(SpecialVariables.This);
if (!isThis || (_context.CurrentTypeDefinitionAst == null && _context.CurrentThisType == null))
// $this is special in script properties and in PowerShell classes
PSTypeName typeName = _context.CurrentThisType ?? new PSTypeName(_context.CurrentTypeDefinitionAst);
inferredTypes.Add(typeName);
return;
}
// Process other well known variables like $true and $pshome
if (SpecialVariables.AllScopeVariables.TryGetValue(astVariablePath.UnqualifiedPath, out Type knownType))
{
if (knownType == typeof(object))
{
if (SpecialVariables.AllScopeVariables.TryGetValue(astVariablePath.UserPath, out Type knownType))
if (_context.TryGetRepresentativeTypeNameFromExpressionSafeEval(variableExpressionAst, out var psType))
{
if (knownType == typeof(object))
{
if (_context.TryGetRepresentativeTypeNameFromExpressionSafeEval(variableExpressionAst, out var psType))
{
inferredTypes.Add(psType);
}
}
else
{
inferredTypes.Add(new PSTypeName(knownType));
}
return;
}
for (int i = 0; i < SpecialVariables.AutomaticVariables.Length; i++)
{
if (!astVariablePath.UserPath.EqualsOrdinalIgnoreCase(SpecialVariables.AutomaticVariables[i]))
{
continue;
}
var type = SpecialVariables.AutomaticVariableTypes[i];
if (type != typeof(object))
{
inferredTypes.Add(new PSTypeName(type));
}
break;
inferredTypes.Add(psType);
}
}
else
{
var typeName = _context.CurrentThisType ?? new PSTypeName(_context.CurrentTypeDefinitionAst);
inferredTypes.Add(typeName);
return;
inferredTypes.Add(new PSTypeName(knownType));
}
}
else
{
inferredTypes.Add(new PSTypeName(_context.CurrentTypeDefinitionAst));
return;
}
if (parent is null)
// Process automatic variables like $MyInvocation and $PSBoundParameters
for (int i = 0; i < SpecialVariables.AutomaticVariables.Length; i++)
{
if (!astVariablePath.UnqualifiedPath.EqualsOrdinalIgnoreCase(SpecialVariables.AutomaticVariables[i]))
{
continue;
}
Type type = SpecialVariables.AutomaticVariableTypes[i];
if (type != typeof(object))
{
inferredTypes.Add(new PSTypeName(type));
}
return;
}
// Look for our variable as a parameter or on the lhs of an assignment - hopefully we'll find either
// a type constraint or at least we can use the rhs to infer the type.
while (parent.Parent != null)
// This visitor + loop finds the start of the current scope and traverses top to bottom to find the nearest variable assignment.
// Then repeats the process for each parent scope.
var assignmentVisitor = new VariableAssignmentVisitor()
{
parent = parent.Parent;
ScopeIsLocal = true,
LocalScopeOnly = variableExpressionAst.VariablePath.IsLocal || variableExpressionAst.VariablePath.IsPrivate,
StopSearchOffset = variableExpressionAst.Extent.StartOffset,
VariableTarget = variableExpressionAst
};
while (currentAst is not null)
{
if (currentAst is IParameterMetadataProvider)
{
currentAst.Visit(assignmentVisitor);
if (assignmentVisitor.LocalScopeOnly
|| assignmentVisitor.LastConstraint is not null
|| ((assignmentVisitor.LastAssignment is not null || assignmentVisitor.LastAssignmentType is not null)
&& (currentAst.Parent is not ScriptBlockExpressionAst scriptBlock || !scriptBlock.IsDotsourced())))
{
// We only care about the parent scopes if no assignment has been made in the current scope
// or if it's a dot sourced scriptblock where an earlier defined type constraint could influence the final type
break;
}
assignmentVisitor.ScopeIsLocal = false;
assignmentVisitor.StopSearchOffset = currentAst.Extent.StartOffset;
}
currentAst = currentAst.Parent;
}
int startOffset = variableExpressionAst.Extent.StartOffset;
var targetAsts = (List<Ast>)AstSearcher.FindAll(
parent,
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)
// The visitor is done finding the last assignment, now we need to infer the type of that assignment.
if (assignmentVisitor.LastConstraint is not null)
{
if (ast is ParameterAst parameterAst)
inferredTypes.Add(new PSTypeName(assignmentVisitor.LastConstraint));
}
else if (assignmentVisitor.LastAssignment is not null)
{
if (assignmentVisitor.EnumerateAssignment)
{
var currentCount = inferredTypes.Count;
inferredTypes.AddRange(InferTypes(parameterAst));
if (inferredTypes.Count != currentCount)
inferredTypes.AddRange(GetInferredEnumeratedTypes(InferTypes(assignmentVisitor.LastAssignment)));
}
else
{
if (assignmentVisitor.LastAssignment is ConvertExpressionAst convertExpression
&& convertExpression.IsRef())
{
return;
if (convertExpression.Parent is InvokeMemberExpressionAst memberInvoke)
{
inferredTypes.AddRange(InferTypeFromRef(memberInvoke, convertExpression));
}
}
else if (assignmentVisitor.RedirectionAssignment && assignmentVisitor.LastAssignment is CommandAst cmdAst)
{
InferTypesFrom(cmdAst, inferredTypes, forRedirection: true);
}
else
{
inferredTypes.AddRange(InferTypes(assignmentVisitor.LastAssignment));
}
}
}
var assignAsts = targetAsts.OfType<AssignmentStatementAst>().ToArray();
// If any of the assignments lhs use a type constraint, then we use that.
// Otherwise, we use the rhs of the "nearest" assignment
for (int i = assignAsts.Length - 1; i >= 0; i--)
else if (assignmentVisitor.LastAssignmentType is not null)
{
if (assignAsts[i].Left is ConvertExpressionAst lhsConvert)
{
inferredTypes.Add(new PSTypeName(lhsConvert.Type.TypeName));
return;
}
}
var foreachAst = targetAsts.OfType<ForEachStatementAst>().FirstOrDefault();
if (foreachAst != null)
{
inferredTypes.AddRange(
GetInferredEnumeratedTypes(InferTypes(foreachAst.Condition)));
return;
}
var commandCompletionAst = targetAsts.OfType<CommandAst>().FirstOrDefault();
if (commandCompletionAst != null)
{
inferredTypes.AddRange(InferTypes(commandCompletionAst));
return;
}
int smallestDiff = int.MaxValue;
AssignmentStatementAst closestAssignment = null;
foreach (var assignAst in assignAsts)
{
var endOffset = assignAst.Extent.EndOffset;
if ((startOffset - endOffset) < smallestDiff)
{
smallestDiff = startOffset - endOffset;
closestAssignment = assignAst;
}
}
if (closestAssignment != null)
{
inferredTypes.AddRange(InferTypes(closestAssignment.Right));
inferredTypes.Add(assignmentVisitor.LastAssignmentType);
}
if (_context.TryGetRepresentativeTypeNameFromExpressionSafeEval(variableExpressionAst, out var evalTypeName))
@@ -2858,6 +2858,399 @@ namespace System.Management.Automation
var i = pipe.PipelineElements.IndexOf(commandAst);
return i != 0 ? pipe.PipelineElements[i - 1] : null;
}
private sealed class VariableAssignmentVisitor : AstVisitor2
{
/// <summary>
/// If set, we only look for local/private assignments in the scope of the variable we are inferring.
/// </summary>
internal bool LocalScopeOnly;
/// <summary>
/// The current scope is local to the variable that is being inferred.
/// </summary>
internal bool ScopeIsLocal;
/// <summary>
/// The variable that we are trying to determine the type of.
/// </summary>
internal VariableExpressionAst VariableTarget;
/// <summary>
/// The last type constraint applied to the variable. This takes priority when determining the type of the variable.
/// </summary>
internal ITypeName LastConstraint;
/// <summary>
/// The last ast that assigned a value to the variable. This determines the value of the variable unless a type constraint has been applied.
/// </summary>
internal Ast LastAssignment;
/// <summary>
/// The inferred type from the most recent assignment. This is only used for stream redirections to variables, or the special OutVariable common parameters.
/// </summary>
internal PSTypeName LastAssignmentType;
/// <summary>
/// Whether or not the types from the last assignment should be enumerated.
/// For assignments made by the PipelineVariable parameter or the foreach statement.
/// </summary>
internal bool EnumerateAssignment;
/// <summary>
/// Whether or not the last assignment was via command redirection.
/// </summary>
internal bool RedirectionAssignment;
internal int StopSearchOffset;
private int LastAssignmentOffset = -1;
private void SetLastAssignment(Ast ast, bool enumerate = false, bool redirectionAssignment = false)
{
if (LastAssignmentOffset < ast.Extent.StartOffset)
{
ClearAssignmentData();
LastAssignment = ast;
EnumerateAssignment = enumerate;
RedirectionAssignment = redirectionAssignment;
LastAssignmentOffset = ast.Extent.StartOffset;
}
}
private void SetLastAssignmentType(PSTypeName typeName, int assignmentOffset)
{
if (LastAssignmentOffset < assignmentOffset)
{
ClearAssignmentData();
LastAssignmentType = typeName;
LastAssignmentOffset = assignmentOffset;
}
}
private void ClearAssignmentData()
{
LastAssignment = null;
LastAssignmentType = null;
EnumerateAssignment = false;
RedirectionAssignment = false;
}
private bool AssignsToTargetVar(VariableExpressionAst foundVar)
{
if (!foundVar.VariablePath.UnqualifiedPath.EqualsOrdinalIgnoreCase(VariableTarget.VariablePath.UnqualifiedPath))
{
return false;
}
int scopeIndex = foundVar.VariablePath.UserPath.IndexOf(':');
string scopeName = scopeIndex == -1 ? string.Empty : foundVar.VariablePath.UserPath.Remove(scopeIndex);
return AssignsToTargetScope(scopeName);
}
private bool AssignsToTargetVar(string userPath)
{
if (string.IsNullOrEmpty(userPath))
{
return false;
}
string scopeName;
string varName;
int scopeIndex = userPath.IndexOf(':');
if (scopeIndex == -1)
{
scopeName = string.Empty;
varName = userPath;
}
else
{
scopeName = userPath.Remove(scopeIndex);
varName = userPath.Substring(scopeIndex + 1);
}
if (!varName.EqualsOrdinalIgnoreCase(VariableTarget.VariablePath.UnqualifiedPath))
{
return false;
}
return AssignsToTargetScope(scopeName);
}
private bool AssignsToTargetScope(string scopeName)
=> LocalScopeOnly
? string.IsNullOrEmpty(scopeName) || scopeName.EqualsOrdinalIgnoreCase("Local") || scopeName.EqualsOrdinalIgnoreCase("Private")
: ScopeIsLocal || !(scopeName.EqualsOrdinalIgnoreCase("Local") || scopeName.EqualsOrdinalIgnoreCase("Private"));
public override AstVisitAction DefaultVisit(Ast ast)
{
if (ast.Extent.StartOffset >= StopSearchOffset)
{
// When visiting do while/until statements, the condition will be visited before the statement block
// The condition itself may not be interesting if it's after the cursor, but the statement block could be
// Example:
// do
// {
// $Var = gci
// $Var.<Tab>
// }
// until($false)
return ast is PipelineBaseAst && ast.Parent is DoUntilStatementAst or DoWhileStatementAst
? AstVisitAction.SkipChildren
: AstVisitAction.StopVisit;
}
return AstVisitAction.Continue;
}
public override AstVisitAction VisitAssignmentStatement(AssignmentStatementAst assignmentStatementAst)
{
if (assignmentStatementAst.Extent.StartOffset >= StopSearchOffset)
{
return assignmentStatementAst.Parent is DoUntilStatementAst or DoWhileStatementAst
? AstVisitAction.SkipChildren
: AstVisitAction.StopVisit;
}
if (assignmentStatementAst.Left is AttributedExpressionAst attributedExpression)
{
var firstConvertExpression = attributedExpression as ConvertExpressionAst;
ExpressionAst child = attributedExpression.Child;
while (child is AttributedExpressionAst attributeChild)
{
if (firstConvertExpression is null && attributeChild is ConvertExpressionAst convertExpression)
{
// Multiple type constraint can be set on a variable like this: [int] [string] $Var1 = 1
// But it's the left most type constraint that determines the final type.
firstConvertExpression = convertExpression;
}
child = attributeChild.Child;
}
if (child is VariableExpressionAst variableExpression && AssignsToTargetVar(variableExpression))
{
if (firstConvertExpression is not null)
{
LastConstraint = firstConvertExpression.Type.TypeName;
}
else
{
SetLastAssignment(assignmentStatementAst.Right);
}
}
}
else if (assignmentStatementAst.Left is VariableExpressionAst variableExpression && AssignsToTargetVar(variableExpression))
{
SetLastAssignment(assignmentStatementAst.Right);
}
return AstVisitAction.Continue;
}
public override AstVisitAction VisitCommand(CommandAst commandAst)
{
if (commandAst.Extent.StartOffset >= StopSearchOffset)
{
return AstVisitAction.StopVisit;
}
string commandName = commandAst.GetCommandName();
if (commandName is not null && CompletionCompleters.s_varModificationCommands.Contains(commandName))
{
StaticBindingResult bindingResult = StaticParameterBinder.BindCommand(commandAst, resolve: false, CompletionCompleters.s_varModificationParameters);
if (bindingResult is not null
&& bindingResult.BoundParameters.TryGetValue("Name", out ParameterBindingResult variableName)
&& variableName.ConstantValue is string nameValue
&& AssignsToTargetVar(nameValue)
&& bindingResult.BoundParameters.TryGetValue("Value", out ParameterBindingResult variableValue))
{
SetLastAssignment(variableValue.Value);
return AstVisitAction.Continue;
}
}
StaticBindingResult bindResult = StaticParameterBinder.BindCommand(commandAst, resolve: false);
if (bindResult is not null)
{
foreach (string parameterName in CompletionCompleters.s_outVarParameters)
{
if (bindResult.BoundParameters.TryGetValue(parameterName, out ParameterBindingResult outVarBind)
&& outVarBind.ConstantValue is string varName
&& AssignsToTargetVar(varName))
{
// The *Variable parameters actually always results in an ArrayList
// But to make type inference of individual elements better, we say it's a generic list.
switch (parameterName)
{
case "ErrorVariable":
case "ev":
SetLastAssignmentType(new PSTypeName(typeof(List<ErrorRecord>)), commandAst.Extent.StartOffset);
break;
case "WarningVariable":
case "wv":
SetLastAssignmentType(new PSTypeName(typeof(List<WarningRecord>)), commandAst.Extent.StartOffset);
break;
case "InformationVariable":
case "iv":
SetLastAssignmentType(new PSTypeName(typeof(List<InformationalRecord>)), commandAst.Extent.StartOffset);
break;
case "OutVariable":
case "ov":
SetLastAssignment(commandAst);
break;
default:
break;
}
return AstVisitAction.Continue;
}
}
if (commandAst.Parent is PipelineAst pipeline && pipeline.Extent.EndOffset > VariableTarget.Extent.StartOffset)
{
foreach (string parameterName in CompletionCompleters.s_pipelineVariableParameters)
{
if (bindResult.BoundParameters.TryGetValue(parameterName, out ParameterBindingResult pipeVarBind)
&& pipeVarBind.ConstantValue is string varName
&& AssignsToTargetVar(varName))
{
SetLastAssignment(commandAst, enumerate: true);
return AstVisitAction.Continue;
}
}
}
}
foreach (RedirectionAst redirection in commandAst.Redirections)
{
if (redirection is FileRedirectionAst fileRedirection
&& fileRedirection.Location is StringConstantExpressionAst redirectTarget
&& redirectTarget.Value.StartsWith("variable:", StringComparison.OrdinalIgnoreCase)
&& redirectTarget.Value.Length > "variable:".Length)
{
string varName = redirectTarget.Value.Substring("variable:".Length);
if (!AssignsToTargetVar(varName))
{
continue;
}
switch (fileRedirection.FromStream)
{
case RedirectionStream.Error:
SetLastAssignmentType(new PSTypeName(typeof(ErrorRecord)), commandAst.Extent.StartOffset);
break;
case RedirectionStream.Warning:
SetLastAssignmentType(new PSTypeName(typeof(WarningRecord)), commandAst.Extent.StartOffset);
break;
case RedirectionStream.Verbose:
SetLastAssignmentType(new PSTypeName(typeof(VerboseRecord)), commandAst.Extent.StartOffset);
break;
case RedirectionStream.Debug:
SetLastAssignmentType(new PSTypeName(typeof(DebugRecord)), commandAst.Extent.StartOffset);
break;
case RedirectionStream.Information:
SetLastAssignmentType(new PSTypeName(typeof(InformationRecord)), commandAst.Extent.StartOffset);
break;
default:
SetLastAssignment(commandAst, redirectionAssignment: true);
break;
}
}
}
return AstVisitAction.Continue;
}
public override AstVisitAction VisitParameter(ParameterAst parameterAst)
{
if (parameterAst.Extent.StartOffset >= StopSearchOffset)
{
return AstVisitAction.StopVisit;
}
if (AssignsToTargetVar(parameterAst.Name))
{
foreach (AttributeBaseAst attribute in parameterAst.Attributes)
{
if (attribute is TypeConstraintAst typeConstraint)
{
LastConstraint = typeConstraint.TypeName;
return AstVisitAction.Continue;
}
}
}
return AstVisitAction.Continue;
}
public override AstVisitAction VisitForEachStatement(ForEachStatementAst forEachStatementAst)
{
if (forEachStatementAst.Extent.StartOffset >= StopSearchOffset)
{
return AstVisitAction.StopVisit;
}
if (AssignsToTargetVar(forEachStatementAst.Variable) && forEachStatementAst.Condition.Extent.EndOffset < VariableTarget.Extent.StartOffset)
{
SetLastAssignment(forEachStatementAst.Condition, enumerate: true);
}
return AstVisitAction.Continue;
}
public override AstVisitAction VisitConvertExpression(ConvertExpressionAst convertExpressionAst)
{
if (convertExpressionAst.IsRef()
&& convertExpressionAst.Child is VariableExpressionAst varAst
&& AssignsToTargetVar(varAst))
{
SetLastAssignment(convertExpressionAst);
}
return AstVisitAction.Continue;
}
public override AstVisitAction VisitAttribute(AttributeAst attributeAst)
{
// Attributes can't assign values to variables so they aren't interesting.
return AstVisitAction.SkipChildren;
}
public override AstVisitAction VisitScriptBlockExpression(ScriptBlockExpressionAst scriptBlockExpressionAst)
{
return scriptBlockExpressionAst.IsDotsourced()
? AstVisitAction.Continue
: AstVisitAction.SkipChildren;
}
public override AstVisitAction VisitDataStatement(DataStatementAst dataStatementAst)
{
if (dataStatementAst.Extent.StartOffset >= StopSearchOffset)
{
return AstVisitAction.StopVisit;
}
if (AssignsToTargetVar(dataStatementAst.Variable) && dataStatementAst.Extent.EndOffset < VariableTarget.Extent.StartOffset)
{
SetLastAssignment(dataStatementAst.Body);
}
return AstVisitAction.SkipChildren;
}
public override AstVisitAction VisitFunctionDefinition(FunctionDefinitionAst functionDefinitionAst)
{
return AstVisitAction.SkipChildren;
}
}
}
internal static class TypeInferenceExtension
@@ -2885,68 +3278,27 @@ namespace System.Management.Automation
return res;
}
public static bool AstAssignsToSameVariable(this VariableExpressionAst variableAst, Ast ast)
public static bool IsDotsourced(this ScriptBlockExpressionAst scriptBlockExpressionAst)
{
var parameterAst = ast as ParameterAst;
var variableAstVariablePath = variableAst.VariablePath;
if (parameterAst != null)
{
return variableAstVariablePath.IsUnscopedVariable &&
parameterAst.Name.VariablePath.UnqualifiedPath.Equals(variableAstVariablePath.UnqualifiedPath, StringComparison.OrdinalIgnoreCase) &&
parameterAst.Parent.Parent.Extent.EndOffset > variableAst.Extent.StartOffset;
}
Ast parent = scriptBlockExpressionAst.Parent;
if (ast is ForEachStatementAst foreachAst)
// This loop checks if the scriptblock is used as a dot sourced command
// or an argument for a command that uses the local scope eg: ForEach-Object -Process {$Var1 = "Hello"}, {Var2 = $true}
while (parent is not null)
{
return variableAstVariablePath.IsUnscopedVariable &&
foreachAst.Variable.VariablePath.UnqualifiedPath.Equals(variableAstVariablePath.UnqualifiedPath, StringComparison.OrdinalIgnoreCase);
}
if (ast is CommandAst commandAst)
{
string[] variableParameters = { "PV", "PipelineVariable", "OV", "OutVariable" };
StaticBindingResult bindingResult = StaticParameterBinder.BindCommand(commandAst, false, variableParameters);
if (bindingResult != null)
if (parent is CommandAst cmdAst)
{
foreach (string commandVariableParameter in variableParameters)
{
if (bindingResult.BoundParameters.TryGetValue(commandVariableParameter, out ParameterBindingResult parameterBindingResult))
{
if (string.Equals(variableAstVariablePath.UnqualifiedPath, (string)parameterBindingResult.ConstantValue, StringComparison.OrdinalIgnoreCase))
{
return true;
}
}
}
string cmdName = cmdAst.GetCommandName();
return CompletionCompleters.s_localScopeCommandNames.Contains(cmdName)
|| (cmdAst.CommandElements[0] is ScriptBlockExpressionAst && cmdAst.InvocationOperator == TokenKind.Dot);
}
return false;
}
if (parent is not CommandExpressionAst and not PipelineAst and not StatementBlockAst and not ArrayExpressionAst and not ArrayLiteralAst)
{
break;
}
var assignmentAst = (AssignmentStatementAst)ast;
var lhs = assignmentAst.Left;
if (lhs is ConvertExpressionAst convertExpr)
{
lhs = convertExpr.Child;
}
if (lhs is not VariableExpressionAst varExpr)
{
return false;
}
var candidateVarPath = varExpr.VariablePath;
if (candidateVarPath.UserPath.Equals(variableAstVariablePath.UserPath, StringComparison.OrdinalIgnoreCase))
{
return true;
}
// The following condition is making an assumption that at script scope, we didn't use $script:, but in the local scope, we did
// If we are searching anything other than script scope, this is wrong.
if (variableAstVariablePath.IsScript && variableAstVariablePath.UnqualifiedPath.Equals(candidateVarPath.UnqualifiedPath, StringComparison.OrdinalIgnoreCase))
{
return true;
parent = parent.Parent;
}
return false;
@@ -1481,6 +1481,30 @@ Describe "Type inference Tests" -tags "CI" {
$res.Name | Should -Be 'System.Management.Automation.Internal.Host.InternalHost'
}
It 'Infers type of variable assigned inside do while loop' {
$res = [AstTypeInference]::InferTypeOf(({
do
{
$Test = 1
$Test
}
while (1)
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.Int32'
}
It 'Infers type of variable assigned inside do until loop' {
$res = [AstTypeInference]::InferTypeOf(({
do
{
$Test = 1
$Test
}
until ($null = gci)
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst] -and $Ast.VariablePath.UserPath -eq 'Test'}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.Int32'
}
It 'Infers type of external applications' {
$res = [AstTypeInference]::InferTypeOf( { pwsh }.Ast)
$res.Name | Should -Be 'System.String'
@@ -1494,6 +1518,178 @@ Describe "Type inference Tests" -tags "CI" {
$null = [AstTypeInference]::InferTypeOf($FoundAst)
}
It 'Ignores type constraint defined outside of scope' {
$res = [AstTypeInference]::InferTypeOf(({
function Outer
{
[string] $Test = "Hello"
function Inner
{
$Test = 2
$Test
}
}
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.Int32'
}
It 'Considers the type constraint defined outside of scope when dot sourcing' {
$res = [AstTypeInference]::InferTypeOf(({
[string] $Test = "Hello"
. {$Test = 2; $Test}
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.String'
}
It 'Infers type of ref assigned variable' {
$res = [AstTypeInference]::InferTypeOf(({
$MyRefVar = $null
$null = [System.Management.Automation.Language.Parser]::ParseInput("", [ref] $MyRefVar, [ref] $null)
$MyRefVar
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.Management.Automation.Language.Token[]'
}
It 'Infers type of variable assigned with New/Set-Variable' {
$res = [AstTypeInference]::InferTypeOf( {
New-Variable -Name Var1 -Value $true | Out-Null
New-Variable -Name Var2 -Value "Hello" | Out-Null
$Var1
$Var2
}.Ast)
$res[0].Name | Should -Be 'System.Boolean'
$res[1].Name | Should -Be 'System.String'
}
It 'Infers type of variable assigned with <ParameterName> common parameter' -TestCases @(
@{ParameterName = "WarningVariable"; ExpectedType = [List[WarningRecord]]}
@{ParameterName = "wv"; ExpectedType = [List[WarningRecord]]}
@{ParameterName = "ErrorVariable"; ExpectedType = [List[ErrorRecord]]}
@{ParameterName = "ev"; ExpectedType = [List[ErrorRecord]]}
@{ParameterName = "InformationVariable"; ExpectedType = [List[InformationalRecord]]}
@{ParameterName = "iv"; ExpectedType = [List[InformationalRecord]]}
@{ParameterName = "OutVariable"; ExpectedType = [guid]}
@{ParameterName = "ov"; ExpectedType = [guid]}
@{ParameterName = "PipelineVariable"; ExpectedType = [guid]}
@{ParameterName = "pv"; ExpectedType = [guid]}
) -Test {
param($ParameterName, $ExpectedType)
$Ast = [scriptblock]::Create("New-Guid -$ParameterName MyOutVar | % {`$MyOutVar}").Ast.FindAll({
param($Ast)
$Ast -is [Language.VariableExpressionAst]
}, $true) | Select-Object -Last 1
$res = [AstTypeInference]::InferTypeOf($Ast)
$res.Type | Should -Be $ExpectedType
}
It 'Infers type of variable assigned via Data statement' {
$res = [AstTypeInference]::InferTypeOf(({
Data MyDataVar {"Hello"}
$MyDataVar
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.String'
}
It 'Infers type of well known variable with global scope' {
$res = [AstTypeInference]::InferTypeOf({$global:true}.Ast)
$res.Name | Should -Be 'System.Boolean'
}
It 'Infers parameter type from closest parameter' {
$res = [AstTypeInference]::InferTypeOf( ({
param([string]$Param1)
function TestFunction {param([bool]$Param1) $Param1}
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.Boolean'
}
It 'Infers variable type from closest foreach statement' {
$res = [AstTypeInference]::InferTypeOf( ({
foreach ($X in 1..10)
{
$X
}
foreach ($X in New-Guid)
{
$X
}
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.Guid'
}
It 'Infers global variable type in child scope' {
$res = [AstTypeInference]::InferTypeOf( ({
$Global:GlobalTest1 = "Hello"
function TestFunction {$GlobalTest1}
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.String'
}
It 'Does not infer private variable type in child scope' {
$res = [AstTypeInference]::InferTypeOf( ({
$Private:PrivateTest1 = "Hello"
function TestFunction {$PrivateTest1}
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Count | Should -Be 0
}
It 'Infers variable assigned with an attribute' {
$res = [AstTypeInference]::InferTypeOf( ({
[ValidateNotNull()]$ValidatedVar1 = New-Guid
$ValidatedVar1
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.Guid'
}
It 'Infers variable assigned with multiple type constraints' {
$res = [AstTypeInference]::InferTypeOf( ({
[int] [string]$MultiConstraintVar1 = "10"
$MultiConstraintVar1
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$res.Name | Should -Be 'System.Int32'
}
It 'Infers variable assigned by redirection' {
$res = [AstTypeInference]::InferTypeOf( ({
New-Guid *>&1 1>variable:RedirVar1; $RedirVar1
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true) | Select-Object -Last 1 ))
$ExpectedTypeNames = @(
[ErrorRecord].FullName
[WarningRecord].FullName
[VerboseRecord].FullName
[DebugRecord].FullName
[InformationRecord].FullName
[guid].FullName
) -join ';'
$res.Name -join ';' | Should -Be $ExpectedTypeNames
}
It 'Infers variables assigned by redirection from specific streams' {
$VarAsts = [List[Language.Ast]]{
[void](New-Guid 1>variable:RedirSuccess 2>variable:RedirError 3>variable:RedirWarning 4>variable:RedirVerbose 5>variable:RedirDebug 6>variable:RedirInfo)
$RedirSuccess
$RedirError
$RedirWarning
$RedirVerbose
$RedirDebug
$RedirInfo
}.Ast.FindAll({param($Ast) $Ast -is [Language.VariableExpressionAst]}, $true)
$ExpectedTypeNames = @(
[guid].FullName
[ErrorRecord].FullName
[WarningRecord].FullName
[VerboseRecord].FullName
[DebugRecord].FullName
[InformationRecord].FullName
)
for ($i = 0; $i -lt $VarAsts.Count; $i++)
{
$res = [AstTypeInference]::InferTypeOf($VarAsts[$i])
$res.Name | Should -Be $ExpectedTypeNames[$i]
}
}
It 'Should infer output from anonymous function' {
$res = [AstTypeInference]::InferTypeOf( { & {"Hello"} }.Ast)
$res.Name | Should -Be 'System.String'