Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
c_generator.py574 linesDownload Raw Back to pycparser
1# ------------------------------------------------------------------------------2# pycparser: c_generator.py3#4# C code generator from pycparser AST nodes.5#6# Eli Bendersky [https://eli.thegreenplace.net/]7# License: BSD8# ------------------------------------------------------------------------------9from typing import Callable, List, Optional10 11from . import c_ast12 13 14class CGenerator:15    """Uses the same visitor pattern as c_ast.NodeVisitor, but modified to16    return a value from each visit method, using string accumulation in17    generic_visit.18    """19 20    indent_level: int21    reduce_parentheses: bool22 23    def __init__(self, reduce_parentheses: bool = False) -> None:24        """Constructs C-code generator25 26        reduce_parentheses:27            if True, eliminates needless parentheses on binary operators28        """29        # Statements start with indentation of self.indent_level spaces, using30        # the _make_indent method.31        self.indent_level = 032        self.reduce_parentheses = reduce_parentheses33 34    def _make_indent(self) -> str:35        return " " * self.indent_level36 37    def visit(self, node: c_ast.Node) -> str:38        method = "visit_" + node.__class__.__name__39        return getattr(self, method, self.generic_visit)(node)40 41    def generic_visit(self, node: Optional[c_ast.Node]) -> str:42        if node is None:43            return ""44        else:45            return "".join(self.visit(c) for c_name, c in node.children())46 47    def visit_Constant(self, n: c_ast.Constant) -> str:48        return n.value49 50    def visit_ID(self, n: c_ast.ID) -> str:51        return n.name52 53    def visit_Pragma(self, n: c_ast.Pragma) -> str:54        ret = "#pragma"55        if n.string:56            ret += " " + n.string57        return ret58 59    def visit_ArrayRef(self, n: c_ast.ArrayRef) -> str:60        arrref = self._parenthesize_unless_simple(n.name)61        return arrref + "[" + self.visit(n.subscript) + "]"62 63    def visit_StructRef(self, n: c_ast.StructRef) -> str:64        sref = self._parenthesize_unless_simple(n.name)65        return sref + n.type + self.visit(n.field)66 67    def visit_FuncCall(self, n: c_ast.FuncCall) -> str:68        fref = self._parenthesize_unless_simple(n.name)69        args = self.visit(n.args) if n.args is not None else ""70        return fref + "(" + args + ")"71 72    def visit_UnaryOp(self, n: c_ast.UnaryOp) -> str:73        match n.op:74            case "sizeof":75                # Always parenthesize the argument of sizeof since it can be76                # a name.77                return f"sizeof({self.visit(n.expr)})"78            case "p++":79                operand = self._parenthesize_unless_simple(n.expr)80                return f"{operand}++"81            case "p--":82                operand = self._parenthesize_unless_simple(n.expr)83                return f"{operand}--"84            case _:85                operand = self._parenthesize_unless_simple(n.expr)86                return f"{n.op}{operand}"87 88    # Precedence map of binary operators:89    precedence_map = {90        # Should be in sync with c_parser.CParser.precedence91        # Higher numbers are stronger binding92        "||": 0,  # weakest binding93        "&&": 1,94        "|": 2,95        "^": 3,96        "&": 4,97        "==": 5,98        "!=": 5,99        ">": 6,100        ">=": 6,101        "<": 6,102        "<=": 6,103        ">>": 7,104        "<<": 7,105        "+": 8,106        "-": 8,107        "*": 9,108        "/": 9,109        "%": 9,  # strongest binding110    }111 112    def visit_BinaryOp(self, n: c_ast.BinaryOp) -> str:113        # Note: all binary operators are left-to-right associative114        #115        # If `n.left.op` has a stronger or equally binding precedence in116        # comparison to `n.op`, no parenthesis are needed for the left:117        # e.g., `(a*b) + c` is equivalent to `a*b + c`, as well as118        #       `(a+b) - c` is equivalent to `a+b - c` (same precedence).119        # If the left operator is weaker binding than the current, then120        # parentheses are necessary:121        # e.g., `(a+b) * c` is NOT equivalent to `a+b * c`.122        lval_str = self._parenthesize_if(123            n.left,124            lambda d: not (125                self._is_simple_node(d)126                or self.reduce_parentheses127                and isinstance(d, c_ast.BinaryOp)128                and self.precedence_map[d.op] >= self.precedence_map[n.op]129            ),130        )131        # If `n.right.op` has a stronger -but not equal- binding precedence,132        # parenthesis can be omitted on the right:133        # e.g., `a + (b*c)` is equivalent to `a + b*c`.134        # If the right operator is weaker or equally binding, then parentheses135        # are necessary:136        # e.g., `a * (b+c)` is NOT equivalent to `a * b+c` and137        #       `a - (b+c)` is NOT equivalent to `a - b+c` (same precedence).138        rval_str = self._parenthesize_if(139            n.right,140            lambda d: not (141                self._is_simple_node(d)142                or self.reduce_parentheses143                and isinstance(d, c_ast.BinaryOp)144                and self.precedence_map[d.op] > self.precedence_map[n.op]145            ),146        )147        return f"{lval_str} {n.op} {rval_str}"148 149    def visit_Assignment(self, n: c_ast.Assignment) -> str:150        rval_str = self._parenthesize_if(151            n.rvalue, lambda n: isinstance(n, c_ast.Assignment)152        )153        return f"{self.visit(n.lvalue)} {n.op} {rval_str}"154 155    def visit_IdentifierType(self, n: c_ast.IdentifierType) -> str:156        return " ".join(n.names)157 158    def _visit_expr(self, n: c_ast.Node) -> str:159        match n:160            case c_ast.InitList():161                return "{" + self.visit(n) + "}"162            case c_ast.ExprList() | c_ast.Compound():163                return "(" + self.visit(n) + ")"164            case _:165                return self.visit(n)166 167    def visit_Decl(self, n: c_ast.Decl, no_type: bool = False) -> str:168        # no_type is used when a Decl is part of a DeclList, where the type is169        # explicitly only for the first declaration in a list.170        #171        s = n.name if no_type else self._generate_decl(n)172        if n.bitsize:173            s += " : " + self.visit(n.bitsize)174        if n.init:175            s += " = " + self._visit_expr(n.init)176        return s177 178    def visit_DeclList(self, n: c_ast.DeclList) -> str:179        s = self.visit(n.decls[0])180        if len(n.decls) > 1:181            s += ", " + ", ".join(182                self.visit_Decl(decl, no_type=True) for decl in n.decls[1:]183            )184        return s185 186    def visit_Typedef(self, n: c_ast.Typedef) -> str:187        s = ""188        if n.storage:189            s += " ".join(n.storage) + " "190        s += self._generate_type(n.type)191        return s192 193    def visit_Cast(self, n: c_ast.Cast) -> str:194        s = "(" + self._generate_type(n.to_type, emit_declname=False) + ")"195        return s + " " + self._parenthesize_unless_simple(n.expr)196 197    def visit_ExprList(self, n: c_ast.ExprList) -> str:198        visited_subexprs = []199        for expr in n.exprs:200            visited_subexprs.append(self._visit_expr(expr))201        return ", ".join(visited_subexprs)202 203    def visit_InitList(self, n: c_ast.InitList) -> str:204        visited_subexprs = []205        for expr in n.exprs:206            visited_subexprs.append(self._visit_expr(expr))207        return ", ".join(visited_subexprs)208 209    def visit_Enum(self, n: c_ast.Enum) -> str:210        return self._generate_struct_union_enum(n, name="enum")211 212    def visit_Alignas(self, n: c_ast.Alignas) -> str:213        return "_Alignas({})".format(self.visit(n.alignment))214 215    def visit_Enumerator(self, n: c_ast.Enumerator) -> str:216        if not n.value:217            return "{indent}{name},\n".format(218                indent=self._make_indent(),219                name=n.name,220            )221        else:222            return "{indent}{name} = {value},\n".format(223                indent=self._make_indent(),224                name=n.name,225                value=self.visit(n.value),226            )227 228    def visit_FuncDef(self, n: c_ast.FuncDef) -> str:229        decl = self.visit(n.decl)230        self.indent_level = 0231        body = self.visit(n.body)232        if n.param_decls:233            knrdecls = ";\n".join(self.visit(p) for p in n.param_decls)234            return decl + "\n" + knrdecls + ";\n" + body + "\n"235        else:236            return decl + "\n" + body + "\n"237 238    def visit_FileAST(self, n: c_ast.FileAST) -> str:239        s = ""240        for ext in n.ext:241            match ext:242                case c_ast.FuncDef():243                    s += self.visit(ext)244                case c_ast.Pragma():245                    s += self.visit(ext) + "\n"246                case _:247                    s += self.visit(ext) + ";\n"248        return s249 250    def visit_Compound(self, n: c_ast.Compound) -> str:251        s = self._make_indent() + "{\n"252        self.indent_level += 2253        if n.block_items:254            s += "".join(self._generate_stmt(stmt) for stmt in n.block_items)255        self.indent_level -= 2256        s += self._make_indent() + "}\n"257        return s258 259    def visit_CompoundLiteral(self, n: c_ast.CompoundLiteral) -> str:260        return "(" + self.visit(n.type) + "){" + self.visit(n.init) + "}"261 262    def visit_EmptyStatement(self, n: c_ast.EmptyStatement) -> str:263        return ";"264 265    def visit_ParamList(self, n: c_ast.ParamList) -> str:266        return ", ".join(self.visit(param) for param in n.params)267 268    def visit_Return(self, n: c_ast.Return) -> str:269        s = "return"270        if n.expr:271            s += " " + self.visit(n.expr)272        return s + ";"273 274    def visit_Break(self, n: c_ast.Break) -> str:275        return "break;"276 277    def visit_Continue(self, n: c_ast.Continue) -> str:278        return "continue;"279 280    def visit_TernaryOp(self, n: c_ast.TernaryOp) -> str:281        s = "(" + self._visit_expr(n.cond) + ") ? "282        s += "(" + self._visit_expr(n.iftrue) + ") : "283        s += "(" + self._visit_expr(n.iffalse) + ")"284        return s285 286    def visit_If(self, n: c_ast.If) -> str:287        s = "if ("288        if n.cond:289            s += self.visit(n.cond)290        s += ")\n"291        s += self._generate_stmt(n.iftrue, add_indent=True)292        if n.iffalse:293            s += self._make_indent() + "else\n"294            s += self._generate_stmt(n.iffalse, add_indent=True)295        return s296 297    def visit_For(self, n: c_ast.For) -> str:298        s = "for ("299        if n.init:300            s += self.visit(n.init)301        s += ";"302        if n.cond:303            s += " " + self.visit(n.cond)304        s += ";"305        if n.next:306            s += " " + self.visit(n.next)307        s += ")\n"308        s += self._generate_stmt(n.stmt, add_indent=True)309        return s310 311    def visit_While(self, n: c_ast.While) -> str:312        s = "while ("313        if n.cond:314            s += self.visit(n.cond)315        s += ")\n"316        s += self._generate_stmt(n.stmt, add_indent=True)317        return s318 319    def visit_DoWhile(self, n: c_ast.DoWhile) -> str:320        s = "do\n"321        s += self._generate_stmt(n.stmt, add_indent=True)322        s += self._make_indent() + "while ("323        if n.cond:324            s += self.visit(n.cond)325        s += ");"326        return s327 328    def visit_StaticAssert(self, n: c_ast.StaticAssert) -> str:329        s = "_Static_assert("330        s += self.visit(n.cond)331        if n.message:332            s += ","333            s += self.visit(n.message)334        s += ")"335        return s336 337    def visit_Switch(self, n: c_ast.Switch) -> str:338        s = "switch (" + self.visit(n.cond) + ")\n"339        s += self._generate_stmt(n.stmt, add_indent=True)340        return s341 342    def visit_Case(self, n: c_ast.Case) -> str:343        s = "case " + self.visit(n.expr) + ":\n"344        for stmt in n.stmts:345            s += self._generate_stmt(stmt, add_indent=True)346        return s347 348    def visit_Default(self, n: c_ast.Default) -> str:349        s = "default:\n"350        for stmt in n.stmts:351            s += self._generate_stmt(stmt, add_indent=True)352        return s353 354    def visit_Label(self, n: c_ast.Label) -> str:355        return n.name + ":\n" + self._generate_stmt(n.stmt)356 357    def visit_Goto(self, n: c_ast.Goto) -> str:358        return "goto " + n.name + ";"359 360    def visit_EllipsisParam(self, n: c_ast.EllipsisParam) -> str:361        return "..."362 363    def visit_Struct(self, n: c_ast.Struct) -> str:364        return self._generate_struct_union_enum(n, "struct")365 366    def visit_Typename(self, n: c_ast.Typename) -> str:367        return self._generate_type(n.type)368 369    def visit_Union(self, n: c_ast.Union) -> str:370        return self._generate_struct_union_enum(n, "union")371 372    def visit_NamedInitializer(self, n: c_ast.NamedInitializer) -> str:373        s = ""374        for name in n.name:375            if isinstance(name, c_ast.ID):376                s += "." + name.name377            else:378                s += "[" + self.visit(name) + "]"379        s += " = " + self._visit_expr(n.expr)380        return s381 382    def visit_FuncDecl(self, n: c_ast.FuncDecl) -> str:383        return self._generate_type(n)384 385    def visit_ArrayDecl(self, n: c_ast.ArrayDecl) -> str:386        return self._generate_type(n, emit_declname=False)387 388    def visit_TypeDecl(self, n: c_ast.TypeDecl) -> str:389        return self._generate_type(n, emit_declname=False)390 391    def visit_PtrDecl(self, n: c_ast.PtrDecl) -> str:392        return self._generate_type(n, emit_declname=False)393 394    def _generate_struct_union_enum(395        self, n: c_ast.Struct | c_ast.Union | c_ast.Enum, name: str396    ) -> str:397        """Generates code for structs, unions, and enums. name should be398        'struct', 'union', or 'enum'.399        """400        if name in ("struct", "union"):401            assert isinstance(n, (c_ast.Struct, c_ast.Union))402            members = n.decls403            body_function = self._generate_struct_union_body404        else:405            assert name == "enum"406            assert isinstance(n, c_ast.Enum)407            members = None if n.values is None else n.values.enumerators408            body_function = self._generate_enum_body409        s = name + " " + (n.name or "")410        if members is not None:411            # None means no members412            # Empty sequence means an empty list of members413            s += "\n"414            s += self._make_indent()415            self.indent_level += 2416            s += "{\n"417            s += body_function(members)418            self.indent_level -= 2419            s += self._make_indent() + "}"420        return s421 422    def _generate_struct_union_body(self, members: List[c_ast.Node]) -> str:423        return "".join(self._generate_stmt(decl) for decl in members)424 425    def _generate_enum_body(self, members: List[c_ast.Enumerator]) -> str:426        # `[:-2] + '\n'` removes the final `,` from the enumerator list427        return "".join(self.visit(value) for value in members)[:-2] + "\n"428 429    def _generate_stmt(self, n: c_ast.Node, add_indent: bool = False) -> str:430        """Generation from a statement node. This method exists as a wrapper431        for individual visit_* methods to handle different treatment of432        some statements in this context.433        """434        if add_indent:435            self.indent_level += 2436        indent = self._make_indent()437        if add_indent:438            self.indent_level -= 2439 440        match n:441            case (442                c_ast.Decl()443                | c_ast.Assignment()444                | c_ast.Cast()445                | c_ast.UnaryOp()446                | c_ast.BinaryOp()447                | c_ast.TernaryOp()448                | c_ast.FuncCall()449                | c_ast.ArrayRef()450                | c_ast.StructRef()451                | c_ast.Constant()452                | c_ast.ID()453                | c_ast.Typedef()454                | c_ast.ExprList()455            ):456                # These can also appear in an expression context so no semicolon457                # is added to them automatically458                #459                return indent + self.visit(n) + ";\n"460            case c_ast.Compound():461                # No extra indentation required before the opening brace of a462                # compound - because it consists of multiple lines it has to463                # compute its own indentation.464                #465                return self.visit(n)466            case c_ast.If():467                return indent + self.visit(n)468            case _:469                return indent + self.visit(n) + "\n"470 471    def _generate_decl(self, n: c_ast.Decl) -> str:472        """Generation from a Decl node."""473        s = ""474        if n.funcspec:475            s = " ".join(n.funcspec) + " "476        if n.storage:477            s += " ".join(n.storage) + " "478        if n.align:479            s += self.visit(n.align[0]) + " "480        s += self._generate_type(n.type)481        return s482 483    def _generate_type(484        self,485        n: c_ast.Node,486        modifiers: List[c_ast.Node] = [],487        emit_declname: bool = True,488    ) -> str:489        """Recursive generation from a type node. n is the type node.490        modifiers collects the PtrDecl, ArrayDecl and FuncDecl modifiers491        encountered on the way down to a TypeDecl, to allow proper492        generation from it.493        """494        # ~ print(n, modifiers)495        match n:496            case c_ast.TypeDecl():497                s = ""498                if n.quals:499                    s += " ".join(n.quals) + " "500                s += self.visit(n.type)501 502                nstr = n.declname if n.declname and emit_declname else ""503                # Resolve modifiers.504                # Wrap in parens to distinguish pointer to array and pointer to505                # function syntax.506                #507                for i, modifier in enumerate(modifiers):508                    match modifier:509                        case c_ast.ArrayDecl():510                            if i != 0 and isinstance(modifiers[i - 1], c_ast.PtrDecl):511                                nstr = "(" + nstr + ")"512                            nstr += "["513                            if modifier.dim_quals:514                                nstr += " ".join(modifier.dim_quals) + " "515                            if modifier.dim is not None:516                                nstr += self.visit(modifier.dim)517                            nstr += "]"518                        case c_ast.FuncDecl():519                            if i != 0 and isinstance(modifiers[i - 1], c_ast.PtrDecl):520                                nstr = "(" + nstr + ")"521                            args = (522                                self.visit(modifier.args)523                                if modifier.args is not None524                                else ""525                            )526                            nstr += "(" + args + ")"527                        case c_ast.PtrDecl():528                            if modifier.quals:529                                quals = " ".join(modifier.quals)530                                suffix = f" {nstr}" if nstr else ""531                                nstr = f"* {quals}{suffix}"532                            else:533                                nstr = "*" + nstr534                if nstr:535                    s += " " + nstr536                return s537            case c_ast.Decl():538                return self._generate_decl(n.type)539            case c_ast.Typename():540                return self._generate_type(n.type, emit_declname=emit_declname)541            case c_ast.IdentifierType():542                return " ".join(n.names) + " "543            case c_ast.ArrayDecl() | c_ast.PtrDecl() | c_ast.FuncDecl():544                return self._generate_type(545                    n.type, modifiers + [n], emit_declname=emit_declname546                )547            case _:548                return self.visit(n)549 550    def _parenthesize_if(551        self, n: c_ast.Node, condition: Callable[[c_ast.Node], bool]552    ) -> str:553        """Visits 'n' and returns its string representation, parenthesized554        if the condition function applied to the node returns True.555        """556        s = self._visit_expr(n)557        if condition(n):558            return "(" + s + ")"559        else:560            return s561 562    def _parenthesize_unless_simple(self, n: c_ast.Node) -> str:563        """Common use case for _parenthesize_if"""564        return self._parenthesize_if(n, lambda d: not self._is_simple_node(d))565 566    def _is_simple_node(self, n: c_ast.Node) -> bool:567        """Returns True for nodes that are "simple" - i.e. nodes that always568        have higher precedence than operators.569        """570        return isinstance(571            n,572            (c_ast.Constant, c_ast.ID, c_ast.ArrayRef, c_ast.StructRef, c_ast.FuncCall),573        )574 
codekingpro/portable-devtools · Team Ai