Files
TransPyC/includes/ast/base.py
2026-07-18 19:25:40 +08:00

475 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import t, c
from stdint import *
import memhub
import string
import viperlib
from .tokens import TokOp
# ============================================================
# ASTKind - 节点类型枚举kind() 虚函数返回值)
# ============================================================
class ASTKind(t.CEnum):
# 核心 15 节点
Module: t.State
FunctionDef: t.State
ClassDef: t.State
Assign: t.State
If: t.State
For: t.State
While: t.State
Return: t.State
Expr: t.State
Name: t.State
Constant: t.State
BinOp: t.State
UnaryOp: t.State
Call: t.State
Compare: t.State
# 模块层
Expression: t.State
Interactive: t.State
FunctionType: t.State
# 语句
Delete: t.State
AugAssign: t.State
AnnAssign: t.State
With: t.State
Raise: t.State
Try: t.State
Assert: t.State
Global: t.State
Nonlocal: t.State
Pass: t.State
Break: t.State
Continue: t.State
Import: t.State
ImportFrom: t.State
Match: t.State
# 表达式
BoolOp: t.State
Lambda: t.State
IfExp: t.State
Dict: t.State
Set: t.State
ListComp: t.State
SetComp: t.State
DictComp: t.State
GeneratorExp: t.State
Await: t.State
Yield: t.State
YieldFrom: t.State
FormattedValue: t.State
JoinedStr: t.State
Attribute: t.State
Subscript: t.State
Starred: t.State
List: t.State
Tuple: t.State
Slice: t.State
NamedExpr: t.State
# 辅助节点
ExceptHandler: t.State
Arguments: t.State
Arg: t.State
Keyword: t.State
Alias: t.State
WithItem: t.State
Comprehension: t.State
OpNode: t.State
# 模式匹配
MatchCase: t.State
MatchValue: t.State
MatchSingleton: t.State
MatchSequence: t.State
MatchMapping: t.State
MatchClass: t.State
MatchStar: t.State
MatchAs: t.State
MatchOr: t.State
# ============================================================
# ASTCtx - 名称上下文Load/Store/Del
# ============================================================
class ASTCtx(t.CEnum):
Load: t.State
Store: t.State
Del: t.State
# ============================================================
# OpKind - 运算符枚举BinOp/UnaryOp/Compare 共用)
# ============================================================
class OpKind(t.CEnum):
# 二元算术/位运算
Add: t.State
Sub: t.State
Mult: t.State
MatMult: t.State
Div: t.State
Mod: t.State
Pow: t.State
LShift: t.State
RShift: t.State
BitOr: t.State
BitXor: t.State
BitAnd: t.State
FloorDiv: t.State
# 布尔运算
And: t.State
Or: t.State
# 比较运算
Eq: t.State
Ne: t.State
Lt: t.State
Le: t.State
Gt: t.State
Ge: t.State
Is: t.State
IsNot: t.State
In: t.State
NotIn: t.State
# 一元运算
Not: t.State
UAdd: t.State
USub: t.State
Invert: t.State
# 空
NoneOp: t.State
# Constant 子类型标识
CONST_INT: t.CDefine = 1
CONST_FLOAT: t.CDefine = 2
CONST_STR: t.CDefine = 3
CONST_BOOL: t.CDefine = 4
CONST_NONE: t.CDefine = 5
# 标志位
FLAG_IS_ASYNC: t.CDefine = 1
FLAG_SIMPLE: t.CDefine = 2
FLAG_HAS_STAR: t.CDefine = 4
# ============================================================
# ASTFlag - 节点标志位枚举
# ============================================================
class ASTFlag(t.CEnum):
IsAsync: t.State
Simple: t.State
HasStar: t.State
# ============================================================
# 辅助函数
# ============================================================
def _copy_str(pool: memhub.MemManager | t.CPtr, src: str) -> str:
"""从 pool 分配并复制 C 字符串"""
if src is None: return None
slen: t.CSizeT = string.strlen(src)
buf: str = pool.alloc(slen + 1)
if buf is None: return None
string.memcpy(buf, src, slen + 1)
return buf
def _set_pos(node: AST | t.CPtr, lineno: t.CInt, col_offset: t.CInt,
end_lineno: t.CInt, end_col_offset: t.CInt):
if node is None: return
node.lineno = lineno
node.col_offset = col_offset
node.end_lineno = end_lineno
node.end_col_offset = end_col_offset
def _inherit_pos(node: AST | t.CPtr, ref: AST | t.CPtr):
if node is None or ref is None: return
node.lineno = ref.lineno
node.col_offset = ref.col_offset
node.end_lineno = ref.end_lineno
node.end_col_offset = ref.end_col_offset
def _init_ast(node: AST | t.CPtr, pool: memhub.MemManager | t.CPtr,
lineno: t.CInt, col: t.CInt):
"""初始化 AST 基类字段"""
node.parent = None
node.pool = pool
node.lineno = lineno
node.col_offset = col
node.end_lineno = lineno
node.end_col_offset = col
node.children = None
def _set_parent_list(lst: list[AST | t.CPtr] | t.CPtr, parent: AST | t.CPtr):
"""为 list 中的每个子节点设置 parent 指针"""
if lst is None: return
n: t.CSizeT = lst.__len__()
i: t.CSizeT = 0
while i < n:
child: AST | t.CPtr = lst.get(i)
if child is not None:
child.parent = parent
i += 1
def _append_child(lst: list[AST | t.CPtr] | t.CPtr, node: AST | t.CPtr,
pool: memhub.MemManager | t.CPtr):
"""向 list 追加子节点(模块级函数,避免与 AST.append 方法名冲突)。
在 AST.append 方法中调用此函数,绕过编译器把 list.append 误解析为
AST.append 虚函数调用的类型推断问题。"""
if lst is None or node is None:
return
lst.append(node)
# ============================================================
# 运算符映射TokOp -> OpKind
# ============================================================
def _binop_from_op(op_id: t.CInt) -> t.CInt:
if op_id == TokOp.Plus: return OpKind.Add
if op_id == TokOp.Minus: return OpKind.Sub
if op_id == TokOp.Star: return OpKind.Mult
if op_id == TokOp.At: return OpKind.MatMult
if op_id == TokOp.Slash: return OpKind.Div
if op_id == TokOp.Percent: return OpKind.Mod
if op_id == TokOp.StarStar: return OpKind.Pow
if op_id == TokOp.LtLt: return OpKind.LShift
if op_id == TokOp.GtGt: return OpKind.RShift
if op_id == TokOp.VBar: return OpKind.BitOr
if op_id == TokOp.Caret: return OpKind.BitXor
if op_id == TokOp.Amp: return OpKind.BitAnd
if op_id == TokOp.DSlash: return OpKind.FloorDiv
return OpKind.NoneOp
def _augop_from_op(op_id: t.CInt) -> t.CInt:
if op_id == TokOp.PlusEq: return OpKind.Add
if op_id == TokOp.MinusEq: return OpKind.Sub
if op_id == TokOp.StarEq: return OpKind.Mult
if op_id == TokOp.AtEq: return OpKind.MatMult
if op_id == TokOp.SlashEq: return OpKind.Div
if op_id == TokOp.PercentEq: return OpKind.Mod
if op_id == TokOp.StarEqEq: return OpKind.Pow
if op_id == TokOp.LtLtEq: return OpKind.LShift
if op_id == TokOp.GtGtEq: return OpKind.RShift
if op_id == TokOp.VBarEq: return OpKind.BitOr
if op_id == TokOp.CaretEq: return OpKind.BitXor
if op_id == TokOp.AmpEq: return OpKind.BitAnd
if op_id == TokOp.DSlashEq: return OpKind.FloorDiv
return OpKind.NoneOp
def _cmpop_from_op(op_id: t.CInt) -> t.CInt:
if op_id == TokOp.EqEq: return OpKind.Eq
if op_id == TokOp.ExclaimEq: return OpKind.Ne
if op_id == TokOp.Less: return OpKind.Lt
if op_id == TokOp.LessEq: return OpKind.Le
if op_id == TokOp.Greater: return OpKind.Gt
if op_id == TokOp.GreaterEq: return OpKind.Ge
return OpKind.NoneOp
# ============================================================
# dump 辅助函数
# ============================================================
def _emit(buf: t.CChar | t.CPtr, size: t.CSizeT, pos: t.CSizeT,
text: str) -> t.CSizeT:
"""追加字面量到 buf+pos返回新 pos"""
if buf is None or pos >= size:
return pos
cur: t.CChar | t.CPtr = (t.CVoid | t.CPtr)(t.CUInt64T(buf) + pos)
rem: t.CSizeT = size - pos
n: t.CInt = viperlib.snprintf(cur, rem, "%s", text)
if n < 0:
return pos
return pos + t.CSizeT(n)
def _emit_str(buf: t.CChar | t.CPtr, size: t.CSizeT, pos: t.CSizeT,
text: str) -> t.CSizeT:
"""追加字符串值到 buf+pos%s 格式)"""
if buf is None or pos >= size:
return pos
cur: t.CChar | t.CPtr = (t.CVoid | t.CPtr)(t.CUInt64T(buf) + pos)
rem: t.CSizeT = size - pos
n: t.CInt = viperlib.snprintf(cur, rem, "%s", text)
if n < 0:
return pos
return pos + t.CSizeT(n)
def _emit_int(buf: t.CChar | t.CPtr, size: t.CSizeT, pos: t.CSizeT,
val: t.CInt64T) -> t.CSizeT:
"""追加整数值到 buf+pos"""
if buf is None or pos >= size:
return pos
cur: t.CChar | t.CPtr = (t.CVoid | t.CPtr)(t.CUInt64T(buf) + pos)
rem: t.CSizeT = size - pos
n: t.CInt = viperlib.snprintf(cur, rem, "%lld", val)
if n < 0:
return pos
return pos + t.CSizeT(n)
def _op_name(op: t.CInt) -> str:
"""返回 OpKind 的字符串名称"""
if op == OpKind.Add: return "Add"
if op == OpKind.Sub: return "Sub"
if op == OpKind.Mult: return "Mult"
if op == OpKind.MatMult: return "MatMult"
if op == OpKind.Div: return "Div"
if op == OpKind.Mod: return "Mod"
if op == OpKind.Pow: return "Pow"
if op == OpKind.LShift: return "LShift"
if op == OpKind.RShift: return "RShift"
if op == OpKind.BitOr: return "BitOr"
if op == OpKind.BitXor: return "BitXor"
if op == OpKind.BitAnd: return "BitAnd"
if op == OpKind.FloorDiv: return "FloorDiv"
if op == OpKind.And: return "And"
if op == OpKind.Or: return "Or"
if op == OpKind.Eq: return "Eq"
if op == OpKind.Ne: return "Ne"
if op == OpKind.Lt: return "Lt"
if op == OpKind.Le: return "Le"
if op == OpKind.Gt: return "Gt"
if op == OpKind.Ge: return "Ge"
if op == OpKind.Is: return "Is"
if op == OpKind.IsNot: return "IsNot"
if op == OpKind.In: return "In"
if op == OpKind.NotIn: return "NotIn"
if op == OpKind.Not: return "Not"
if op == OpKind.UAdd: return "UAdd"
if op == OpKind.USub: return "USub"
if op == OpKind.Invert: return "Invert"
return "NoneOp"
def _dump_list(lst: list[AST | t.CPtr] | t.CPtr, buf: t.CChar | t.CPtr,
size: t.CSizeT, pos: t.CSizeT) -> t.CSizeT:
"""遍历 list[AST | CPtr] 容器,多态 dump 每个元素"""
if lst is None:
pos = _emit(buf, size, pos, "[]")
return pos
pos = _emit(buf, size, pos, "[")
n: t.CSizeT = lst.__len__()
i: t.CSizeT = 0
while i < n:
if i > 0:
pos = _emit(buf, size, pos, ", ")
child: AST | t.CPtr = lst.get(i)
if child is not None:
pos = child.dump(buf, size, pos)
i += 1
pos = _emit(buf, size, pos, "]")
return pos
def _dump_op_list(lst: t.CPtr, buf: t.CChar | t.CPtr,
size: t.CSizeT, pos: t.CSizeT) -> t.CSizeT:
"""遍历 list[CInt] (OpKind 值),输出 op 名称列表"""
if lst is None:
pos = _emit(buf, size, pos, "[]")
return pos
ops_list: list[t.CInt] | t.CPtr = (list[t.CInt] | t.CPtr)(lst)
pos = _emit(buf, size, pos, "[")
n: t.CSizeT = ops_list.__len__()
i: t.CSizeT = 0
while i < n:
if i > 0:
pos = _emit(buf, size, pos, ", ")
op: t.CInt = ops_list.get(i)
pos = _emit(buf, size, pos, _op_name(op))
i += 1
pos = _emit(buf, size, pos, "]")
return pos
# ============================================================
# AST - 多态基类
#
# @t.CVTable 启用 vtable支持 kind()/type_name()/dump()/accept() 虚函数
# 子节点用 list[AST | CPtr] 容器存储O(1) 随机访问,比链表高效)
# parent 字段维护父指针(用于向上遍历)
# ============================================================
@t.CVTable
class AST:
"""AST 节点基类。所有具体节点类继承此类。
字段:
parent: 父节点指针(向上遍历)
pool: 分配器(用于子节点/字符串分配)
lineno/col_offset: 起始位置
end_lineno/end_col_offset: 结束位置
children: 子节点列表body 语句块,通过 append 添加)
"""
parent: AST | t.CPtr
pool: memhub.MemManager | t.CPtr
lineno: t.CInt
col_offset: t.CInt
end_lineno: t.CInt
end_col_offset: t.CInt
children: list[AST | t.CPtr] | t.CPtr
def kind(self) -> t.CInt:
"""返回节点类型ASTKind 值),子类覆盖"""
return 0
def type_name(self) -> str:
"""返回节点类型名,子类覆盖"""
return "AST"
def dump(self, buf: t.CChar | t.CPtr, size: t.CSizeT,
pos: t.CSizeT) -> t.CSizeT:
"""多态 dump子类覆盖以自定义格式。返回写入后的 pos"""
return pos
def accept(self, visitor: t.CPtr):
"""访问者模式入口待实现ASTVisitor 在 __visitor.py 定义后,
子类覆盖此方法分派到 visitor.visit_X(self)。当前为空实现。"""
pass
def append(self, node: AST | t.CPtr):
"""将 node 追加为子节点(添加到 children 列表),设置 parent 指针。
用于 parser 构建语句块module.append(stmt), if_node.append(stmt) 等。
首次调用时懒初始化 children 列表。
注意list[AST|CPtr] 当前按值复制存储(编译器限制),所以必须先设置
node.parent 再 append这样副本中的 parent 才是正确的。"""
if self is None or node is None:
return
if self.children is None:
self.children = list[AST | t.CPtr](self.pool, 8)
node.parent = self
_append_child(self.children, node, self.pool)
# ============================================================
# 模块级 dump 函数(供 ast.dump(node, buf, size) 调用)
# ============================================================
def dump(node: AST | t.CPtr, buf: t.CChar | t.CPtr, size: t.CSizeT):
"""将 AST 节点序列化到 buf容量 size"""
if node is None or buf is None or size == 0:
return
buf[0] = '\0'
final_pos: t.CSizeT = node.dump(buf, size, 0)
# NUL 终止,防止 printf 越界
if final_pos < size:
buf[final_pos] = '\0'
else:
buf[size - 1] = '\0'