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'