Files
MycLib/Src/AST/Myc.Ast.Compiler.TCO.pas
T
2025-11-02 19:38:52 +01:00

243 lines
7.2 KiB
ObjectPascal

unit Myc.Ast.Compiler.TCO;
interface
uses
System.SysUtils,
System.Classes,
System.Generics.Collections,
Myc.Data.Value,
Myc.Ast.Nodes,
Myc.Ast.Visitor,
Myc.Ast.Scope,
Myc.Ast.Types,
Myc.Ast;
type
IAstTCO = interface(IAstVisitor)
function Execute(const RootNode: IAstNode): IAstNode;
end;
// This transformer runs *after* the Lowerer (Phase 4).
// Its sole responsibility is to identify tail calls (TCO)
// and set the `IsTailCall` flag on TFunctionCallNode.
TAstTCO = class(TAstTransformer, IAstTCO)
private
FIsTailStack: TStack<Boolean>;
FNextIsTail: Boolean;
protected
function Accept(const Node: IAstNode): IAstNode; override;
// Overrides for TCO propagation
function VisitBlockExpression(const Node: IBlockExpressionNode): IAstNode; override;
function VisitIfExpression(const Node: IIfExpressionNode): IAstNode; override;
function VisitTernaryExpression(const Node: ITernaryExpressionNode): IAstNode; override;
function VisitLambdaExpression(const Node: ILambdaExpressionNode): IAstNode; override;
function VisitRecurNode(const Node: IRecurNode): IAstNode; override;
function VisitMacroExpansionNode(const Node: IMacroExpansionNode): IAstNode; override;
// Operands are never in tail position.
function VisitBinaryExpression(const Node: IBinaryExpressionNode): IAstNode; override;
function VisitUnaryExpression(const Node: IUnaryExpressionNode): IAstNode; override;
// The core TCO logic
function VisitFunctionCall(const Node: IFunctionCallNode): IAstNode; override;
public
constructor Create;
destructor Destroy; override;
function Execute(const RootNode: IAstNode): IAstNode;
class function Optimize(const RootNode: IAstNode): IAstNode; static;
end;
implementation
{ TAstTCO }
constructor TAstTCO.Create;
begin
inherited Create;
FIsTailStack := TStack<Boolean>.Create;
FNextIsTail := True; // The root expression is in tail position
end;
destructor TAstTCO.Destroy;
begin
FIsTailStack.Free;
inherited;
end;
class function TAstTCO.Optimize(const RootNode: IAstNode): IAstNode;
begin
var optimizer := TAstTCO.Create as IAstTCO;
Result := optimizer.Execute(RootNode);
end;
function TAstTCO.Execute(const RootNode: IAstNode): IAstNode;
begin
Result := Accept(RootNode); // Use IAstNode-returning Accept
if not Assigned(Result) then
Result := TAst.Block([]);
end;
function TAstTCO.Accept(const Node: IAstNode): IAstNode;
begin
if (not Assigned(Node)) then
begin
Result := nil;
exit;
end;
FIsTailStack.Push(FNextIsTail);
try
// Call inherited Accept, which handles the data unwrapping
Result := inherited Accept(Node);
finally
FNextIsTail := FIsTailStack.Pop;
end;
end;
function TAstTCO.VisitBinaryExpression(const Node: IBinaryExpressionNode): IAstNode;
begin
// Operands are never in tail position.
FNextIsTail := False;
// Call inherited, which will visit Left and Right with FNextIsTail = False
Result := inherited VisitBinaryExpression(Node);
end;
function TAstTCO.VisitUnaryExpression(const Node: IUnaryExpressionNode): IAstNode;
begin
// Operand is never in tail position.
FNextIsTail := False;
// Call inherited, which will visit Right with FNextIsTail = False
Result := inherited VisitUnaryExpression(Node);
end;
function TAstTCO.VisitBlockExpression(const Node: IBlockExpressionNode): IAstNode;
var
i: Integer;
isContextTail: Boolean;
N: TBlockExpressionNode;
begin
N := (Node as TBlockExpressionNode);
isContextTail := FIsTailStack.Peek;
// We must manually iterate here to set FNextIsTail for each child
for i := 0 to High(N.Expressions) do
begin
// Only the last expression in a block is in tail position
FNextIsTail := isContextTail and (i = High(N.Expressions));
N.Expressions[i] := Accept(N.Expressions[i]);
end;
Result := N;
end;
function TAstTCO.VisitIfExpression(const Node: IIfExpressionNode): IAstNode;
var
isContextTail: Boolean;
N: TIfExpressionNode;
begin
N := (Node as TIfExpressionNode);
isContextTail := FIsTailStack.Peek;
// Condition is never in tail position
FNextIsTail := False;
N.Condition := Accept(N.Condition);
// Then/Else branches ARE in tail position if the IfExpr is
FNextIsTail := isContextTail;
N.ThenBranch := Accept(N.ThenBranch);
N.ElseBranch := Accept(N.ElseBranch);
Result := N;
end;
function TAstTCO.VisitTernaryExpression(const Node: ITernaryExpressionNode): IAstNode;
var
isContextTail: Boolean;
N: TTernaryExpressionNode;
begin
N := (Node as TTernaryExpressionNode);
isContextTail := FIsTailStack.Peek;
// Condition is never in tail position
FNextIsTail := False;
N.Condition := Accept(N.Condition);
// Then/Else branches ARE in tail position if the TernaryExpr is
FNextIsTail := isContextTail;
N.ThenBranch := Accept(N.ThenBranch);
N.ElseBranch := Accept(N.ElseBranch);
Result := N;
end;
function TAstTCO.VisitLambdaExpression(const Node: ILambdaExpressionNode): IAstNode;
var
N: TLambdaExpressionNode;
begin
N := (Node as TLambdaExpressionNode);
// The body of a lambda is *always* a tail position (relative to the lambda)
FNextIsTail := True;
N.Body := Accept(N.Body);
Result := N;
end;
function TAstTCO.VisitRecurNode(const Node: IRecurNode): IAstNode;
var
N: TRecurNode;
begin
if not FIsTailStack.Peek then
raise Exception.Create('''recur'' can only be used in a tail position.');
N := (Node as TRecurNode);
// Arguments are not in tail position
FNextIsTail := False;
N.Arguments := AcceptNodes<IAstNode>(N.Arguments, function(Node: IAstNode): IAstNode begin exit(Node) end);
Result := N;
end;
function TAstTCO.VisitMacroExpansionNode(const Node: IMacroExpansionNode): IAstNode;
var
N: TMacroExpansionNode;
begin
N := (Node as TMacroExpansionNode);
// Propagate tail call status to the expanded body
FNextIsTail := FIsTailStack.Peek;
N.ExpandedBody := Accept(N.ExpandedBody);
// Also visit the original call nodes, though they are not in tail pos
FNextIsTail := False;
N.Callee := Accept(N.Callee);
N.Arguments := AcceptNodes<IAstNode>(N.Arguments, function(Node: IAstNode): IAstNode begin exit(Node) end);
Result := N;
end;
function TAstTCO.VisitFunctionCall(const Node: IFunctionCallNode): IAstNode;
var
isTailCall: Boolean;
N: TFunctionCallNode;
begin
isTailCall := FIsTailStack.Peek;
N := (Node as TFunctionCallNode);
// Arguments are not in tail position
FNextIsTail := False;
N.Callee := Accept(N.Callee);
N.Arguments := AcceptNodes<IAstNode>(N.Arguments, function(Node: IAstNode): IAstNode begin exit(Node) end);
// Mutate the node with the TCO status
N.IsTailCall := isTailCall;
Result := N;
end;
end.