pyir 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- pyir/__init__.py +78 -0
- pyir/__main__.py +47 -0
- pyir/errors.py +38 -0
- pyir/interp.py +110 -0
- pyir/ir.py +412 -0
- pyir/lower.py +523 -0
- pyir/passes.py +390 -0
- pyir/printer.py +32 -0
- pyir/py.typed +0 -0
- pyir-0.1.0.dist-info/METADATA +251 -0
- pyir-0.1.0.dist-info/RECORD +15 -0
- pyir-0.1.0.dist-info/WHEEL +5 -0
- pyir-0.1.0.dist-info/entry_points.txt +2 -0
- pyir-0.1.0.dist-info/licenses/LICENSE +21 -0
- pyir-0.1.0.dist-info/top_level.txt +1 -0
pyir/__init__.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
"""pyir - Python Intermediate Representation.
|
|
2
|
+
|
|
3
|
+
Lower a small, statically-tractable subset of Python into a three-address-code
|
|
4
|
+
IR organised in basic blocks, print it, optimise it and interpret it.
|
|
5
|
+
|
|
6
|
+
Typical use::
|
|
7
|
+
|
|
8
|
+
import pyir
|
|
9
|
+
|
|
10
|
+
module = pyir.compile_source(source)
|
|
11
|
+
print(pyir.format_module(module))
|
|
12
|
+
pyir.optimize(module)
|
|
13
|
+
result = pyir.interpret(module, "gcd", 48, 18)
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
from .errors import IRRuntimeError, IRVerificationError, PyIRError, UnsupportedSyntaxError
|
|
19
|
+
from .interp import Interpreter, interpret
|
|
20
|
+
from .ir import (
|
|
21
|
+
BUILTINS,
|
|
22
|
+
BasicBlock,
|
|
23
|
+
BinOp,
|
|
24
|
+
Branch,
|
|
25
|
+
Call,
|
|
26
|
+
Const,
|
|
27
|
+
Copy,
|
|
28
|
+
Function,
|
|
29
|
+
Jump,
|
|
30
|
+
Module,
|
|
31
|
+
Return,
|
|
32
|
+
UnaryOp,
|
|
33
|
+
Var,
|
|
34
|
+
)
|
|
35
|
+
from .lower import compile_function, compile_source
|
|
36
|
+
from .passes import DEFAULT_PASSES, constant_fold, eliminate_dead_code, optimize, simplify_cfg
|
|
37
|
+
from .printer import format_block, format_function, format_module
|
|
38
|
+
|
|
39
|
+
__version__ = "0.1.0"
|
|
40
|
+
|
|
41
|
+
__all__ = [
|
|
42
|
+
"__version__",
|
|
43
|
+
# front end
|
|
44
|
+
"compile_source",
|
|
45
|
+
"compile_function",
|
|
46
|
+
# IR
|
|
47
|
+
"Module",
|
|
48
|
+
"Function",
|
|
49
|
+
"BasicBlock",
|
|
50
|
+
"Const",
|
|
51
|
+
"Var",
|
|
52
|
+
"Copy",
|
|
53
|
+
"BinOp",
|
|
54
|
+
"UnaryOp",
|
|
55
|
+
"Call",
|
|
56
|
+
"Jump",
|
|
57
|
+
"Branch",
|
|
58
|
+
"Return",
|
|
59
|
+
"BUILTINS",
|
|
60
|
+
# printing
|
|
61
|
+
"format_module",
|
|
62
|
+
"format_function",
|
|
63
|
+
"format_block",
|
|
64
|
+
# optimisation
|
|
65
|
+
"optimize",
|
|
66
|
+
"constant_fold",
|
|
67
|
+
"eliminate_dead_code",
|
|
68
|
+
"simplify_cfg",
|
|
69
|
+
"DEFAULT_PASSES",
|
|
70
|
+
# execution
|
|
71
|
+
"interpret",
|
|
72
|
+
"Interpreter",
|
|
73
|
+
# errors
|
|
74
|
+
"PyIRError",
|
|
75
|
+
"UnsupportedSyntaxError",
|
|
76
|
+
"IRVerificationError",
|
|
77
|
+
"IRRuntimeError",
|
|
78
|
+
]
|
pyir/__main__.py
ADDED
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
"""Command line interface: ``python -m pyir FILE [-O] [--run FUNC ARG...]``."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import argparse
|
|
6
|
+
import ast
|
|
7
|
+
import sys
|
|
8
|
+
from typing import Sequence
|
|
9
|
+
|
|
10
|
+
from .errors import PyIRError
|
|
11
|
+
from .interp import interpret
|
|
12
|
+
from .lower import compile_source
|
|
13
|
+
from .passes import optimize
|
|
14
|
+
from .printer import format_module
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def main(argv: Sequence[str] | None = None) -> int:
|
|
18
|
+
"""Entry point for ``python -m pyir`` and the ``pyir`` console script."""
|
|
19
|
+
parser = argparse.ArgumentParser(prog="pyir", description="Lower Python functions to pyir IR.")
|
|
20
|
+
parser.add_argument("file", help="Python source file containing only function definitions ('-' for stdin)")
|
|
21
|
+
parser.add_argument("-O", "--optimize", action="store_true", help="run the default optimisation passes")
|
|
22
|
+
parser.add_argument(
|
|
23
|
+
"--run",
|
|
24
|
+
nargs="+",
|
|
25
|
+
metavar=("FUNC", "ARG"),
|
|
26
|
+
help="interpret FUNC with literal arguments (e.g. --run gcd 48 18) instead of printing IR",
|
|
27
|
+
)
|
|
28
|
+
args = parser.parse_args(argv)
|
|
29
|
+
|
|
30
|
+
source = sys.stdin.read() if args.file == "-" else open(args.file, encoding="utf-8").read()
|
|
31
|
+
try:
|
|
32
|
+
module = compile_source(source)
|
|
33
|
+
if args.optimize:
|
|
34
|
+
optimize(module)
|
|
35
|
+
if args.run:
|
|
36
|
+
func, *raw = args.run
|
|
37
|
+
print(repr(interpret(module, func, *(ast.literal_eval(a) for a in raw))))
|
|
38
|
+
else:
|
|
39
|
+
print(format_module(module))
|
|
40
|
+
except (PyIRError, SyntaxError) as exc:
|
|
41
|
+
print(f"pyir: error: {exc}", file=sys.stderr)
|
|
42
|
+
return 1
|
|
43
|
+
return 0
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
if __name__ == "__main__":
|
|
47
|
+
sys.exit(main())
|
pyir/errors.py
ADDED
|
@@ -0,0 +1,38 @@
|
|
|
1
|
+
"""Exception hierarchy for :mod:`pyir`."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
__all__ = [
|
|
6
|
+
"PyIRError",
|
|
7
|
+
"UnsupportedSyntaxError",
|
|
8
|
+
"IRVerificationError",
|
|
9
|
+
"IRRuntimeError",
|
|
10
|
+
]
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class PyIRError(Exception):
|
|
14
|
+
"""Base class for every error raised by pyir."""
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class UnsupportedSyntaxError(PyIRError):
|
|
18
|
+
"""Raised when the lowering step meets Python syntax outside the supported subset."""
|
|
19
|
+
|
|
20
|
+
def __init__(self, message: str, lineno: int | None = None) -> None:
|
|
21
|
+
self.lineno = lineno
|
|
22
|
+
if lineno is not None:
|
|
23
|
+
message = f"line {lineno}: {message}"
|
|
24
|
+
super().__init__(message)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
class IRVerificationError(PyIRError):
|
|
28
|
+
"""Raised by :meth:`pyir.ir.Function.verify` when a function is malformed."""
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class IRRuntimeError(PyIRError):
|
|
32
|
+
"""Raised by the IR interpreter for IR-level failures.
|
|
33
|
+
|
|
34
|
+
Examples are reading an unassigned variable, calling an unknown function or
|
|
35
|
+
exceeding the step budget. Exceptions raised by the operations themselves
|
|
36
|
+
(``ZeroDivisionError``, ``TypeError``, ...) propagate unchanged so they match
|
|
37
|
+
what the original Python code would raise.
|
|
38
|
+
"""
|
pyir/interp.py
ADDED
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
"""A straightforward interpreter for pyir IR, used to check semantics."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
from .errors import IRRuntimeError
|
|
8
|
+
from .ir import (
|
|
9
|
+
BINARY_OPS,
|
|
10
|
+
BUILTINS,
|
|
11
|
+
UNARY_OPS,
|
|
12
|
+
BinOp,
|
|
13
|
+
Branch,
|
|
14
|
+
Call,
|
|
15
|
+
Const,
|
|
16
|
+
Copy,
|
|
17
|
+
Function,
|
|
18
|
+
Jump,
|
|
19
|
+
Module,
|
|
20
|
+
Operand,
|
|
21
|
+
Return,
|
|
22
|
+
UnaryOp,
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
__all__ = ["Interpreter", "interpret", "DEFAULT_MAX_STEPS"]
|
|
26
|
+
|
|
27
|
+
DEFAULT_MAX_STEPS = 10_000_000
|
|
28
|
+
"""Default budget of executed instructions (terminators included) per :meth:`Interpreter.run`."""
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class Interpreter:
|
|
32
|
+
"""Executes functions of a :class:`~pyir.ir.Module`.
|
|
33
|
+
|
|
34
|
+
Attributes:
|
|
35
|
+
module: The module whose functions can be run and called.
|
|
36
|
+
max_steps: Maximum number of instructions (including terminators) a
|
|
37
|
+
single :meth:`run` may execute, across nested calls.
|
|
38
|
+
steps: Number of instructions executed by the most recent :meth:`run`.
|
|
39
|
+
"""
|
|
40
|
+
|
|
41
|
+
def __init__(self, module: Module, max_steps: int = DEFAULT_MAX_STEPS) -> None:
|
|
42
|
+
self.module = module
|
|
43
|
+
self.max_steps = max_steps
|
|
44
|
+
self.steps = 0
|
|
45
|
+
|
|
46
|
+
def run(self, name: str, *args: Any) -> Any:
|
|
47
|
+
"""Call function ``name`` with ``args`` and return its result.
|
|
48
|
+
|
|
49
|
+
Raises:
|
|
50
|
+
IRRuntimeError: on unknown functions, unassigned variables or when
|
|
51
|
+
the step budget is exhausted.
|
|
52
|
+
TypeError: when the argument count does not match the parameters.
|
|
53
|
+
Exception: anything the executed operations raise
|
|
54
|
+
(e.g. ``ZeroDivisionError``) propagates unchanged.
|
|
55
|
+
"""
|
|
56
|
+
self.steps = 0
|
|
57
|
+
return self._call(name, list(args))
|
|
58
|
+
|
|
59
|
+
def _call(self, name: str, args: list[Any]) -> Any:
|
|
60
|
+
fn = self.module.functions.get(name)
|
|
61
|
+
if fn is not None:
|
|
62
|
+
return self._execute(fn, args)
|
|
63
|
+
builtin = BUILTINS.get(name)
|
|
64
|
+
if builtin is not None:
|
|
65
|
+
return builtin(*args)
|
|
66
|
+
raise IRRuntimeError(f"call to unknown function {name!r}")
|
|
67
|
+
|
|
68
|
+
def _execute(self, fn: Function, args: list[Any]) -> Any:
|
|
69
|
+
if len(args) != len(fn.params):
|
|
70
|
+
raise TypeError(f"{fn.name}() takes {len(fn.params)} positional arguments but {len(args)} were given")
|
|
71
|
+
env: dict[str, Any] = dict(zip(fn.params, args))
|
|
72
|
+
|
|
73
|
+
def read(op: Operand) -> Any:
|
|
74
|
+
if isinstance(op, Const):
|
|
75
|
+
return op.value
|
|
76
|
+
try:
|
|
77
|
+
return env[op.name]
|
|
78
|
+
except KeyError:
|
|
79
|
+
raise IRRuntimeError(f"{fn.name}: variable {op.name!r} read before assignment") from None
|
|
80
|
+
|
|
81
|
+
block = fn.entry
|
|
82
|
+
while True:
|
|
83
|
+
self.steps += len(block.instructions) + 1
|
|
84
|
+
if self.steps > self.max_steps:
|
|
85
|
+
raise IRRuntimeError(f"step budget of {self.max_steps} exhausted in {fn.name}")
|
|
86
|
+
for inst in block.instructions:
|
|
87
|
+
if isinstance(inst, Copy):
|
|
88
|
+
env[inst.dest.name] = read(inst.src)
|
|
89
|
+
elif isinstance(inst, BinOp):
|
|
90
|
+
env[inst.dest.name] = BINARY_OPS[inst.op](read(inst.left), read(inst.right))
|
|
91
|
+
elif isinstance(inst, UnaryOp):
|
|
92
|
+
env[inst.dest.name] = UNARY_OPS[inst.op](read(inst.operand))
|
|
93
|
+
elif isinstance(inst, Call):
|
|
94
|
+
env[inst.dest.name] = self._call(inst.func, [read(a) for a in inst.args])
|
|
95
|
+
else: # pragma: no cover - guarded by the Instruction union
|
|
96
|
+
raise IRRuntimeError(f"unknown instruction {inst!r}")
|
|
97
|
+
term = block.terminator
|
|
98
|
+
if isinstance(term, Return):
|
|
99
|
+
return read(term.value)
|
|
100
|
+
if isinstance(term, Jump):
|
|
101
|
+
block = fn.blocks[term.target]
|
|
102
|
+
elif isinstance(term, Branch):
|
|
103
|
+
block = fn.blocks[term.if_true if read(term.cond) else term.if_false]
|
|
104
|
+
else:
|
|
105
|
+
raise IRRuntimeError(f"{fn.name}: block {block.label!r} has no terminator")
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def interpret(module: Module, name: str, *args: Any, max_steps: int = DEFAULT_MAX_STEPS) -> Any:
|
|
109
|
+
"""Convenience wrapper: ``Interpreter(module, max_steps).run(name, *args)``."""
|
|
110
|
+
return Interpreter(module, max_steps).run(name, *args)
|
pyir/ir.py
ADDED
|
@@ -0,0 +1,412 @@
|
|
|
1
|
+
"""Core IR data structures.
|
|
2
|
+
|
|
3
|
+
The IR is a classic *three-address code* organised into basic blocks:
|
|
4
|
+
|
|
5
|
+
* A :class:`Module` holds named :class:`Function` objects.
|
|
6
|
+
* A :class:`Function` holds an ordered mapping of labels to :class:`BasicBlock`.
|
|
7
|
+
The first block is the entry block.
|
|
8
|
+
* A :class:`BasicBlock` is a list of straight-line instructions followed by
|
|
9
|
+
exactly one terminator (:class:`Jump`, :class:`Branch` or :class:`Return`).
|
|
10
|
+
* Operands are either a :class:`Const` or a :class:`Var`. Variables named
|
|
11
|
+
after Python locals keep their source name; compiler temporaries are
|
|
12
|
+
spelled ``%tN``.
|
|
13
|
+
|
|
14
|
+
Variables may be assigned more than once (the IR is *not* SSA), which keeps
|
|
15
|
+
lowering and interpretation simple; the optimisation passes use classic
|
|
16
|
+
dataflow analyses that do not rely on single assignment.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
from __future__ import annotations
|
|
20
|
+
|
|
21
|
+
import operator
|
|
22
|
+
from dataclasses import dataclass, field
|
|
23
|
+
from typing import Any, Callable, Iterator, Union
|
|
24
|
+
|
|
25
|
+
from .errors import IRVerificationError
|
|
26
|
+
|
|
27
|
+
__all__ = [
|
|
28
|
+
"Const",
|
|
29
|
+
"Var",
|
|
30
|
+
"Operand",
|
|
31
|
+
"BINARY_OPS",
|
|
32
|
+
"UNARY_OPS",
|
|
33
|
+
"BUILTINS",
|
|
34
|
+
"Copy",
|
|
35
|
+
"BinOp",
|
|
36
|
+
"UnaryOp",
|
|
37
|
+
"Call",
|
|
38
|
+
"Instruction",
|
|
39
|
+
"Jump",
|
|
40
|
+
"Branch",
|
|
41
|
+
"Return",
|
|
42
|
+
"Terminator",
|
|
43
|
+
"BasicBlock",
|
|
44
|
+
"Function",
|
|
45
|
+
"Module",
|
|
46
|
+
]
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
# --------------------------------------------------------------------------
|
|
50
|
+
# Operands
|
|
51
|
+
# --------------------------------------------------------------------------
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
class Const:
|
|
55
|
+
"""A literal constant operand (``int``, ``float``, ``bool``, ``str`` or ``None``).
|
|
56
|
+
|
|
57
|
+
Equality is type-aware so that ``Const(1)``, ``Const(True)`` and
|
|
58
|
+
``Const(1.0)`` are all distinct, and ``Const(0.0) != Const(-0.0)``.
|
|
59
|
+
"""
|
|
60
|
+
|
|
61
|
+
__slots__ = ("value",)
|
|
62
|
+
|
|
63
|
+
def __init__(self, value: Any) -> None:
|
|
64
|
+
self.value = value
|
|
65
|
+
|
|
66
|
+
def _key(self) -> tuple[type, str]:
|
|
67
|
+
return (type(self.value), repr(self.value))
|
|
68
|
+
|
|
69
|
+
def __eq__(self, other: object) -> bool:
|
|
70
|
+
return isinstance(other, Const) and self._key() == other._key()
|
|
71
|
+
|
|
72
|
+
def __hash__(self) -> int:
|
|
73
|
+
return hash(self._key())
|
|
74
|
+
|
|
75
|
+
def __repr__(self) -> str:
|
|
76
|
+
return f"Const({self.value!r})"
|
|
77
|
+
|
|
78
|
+
def __str__(self) -> str:
|
|
79
|
+
return repr(self.value)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
@dataclass(frozen=True)
|
|
83
|
+
class Var:
|
|
84
|
+
"""A named variable: a Python local (``x``) or a compiler temporary (``%t3``)."""
|
|
85
|
+
|
|
86
|
+
name: str
|
|
87
|
+
|
|
88
|
+
@property
|
|
89
|
+
def is_temp(self) -> bool:
|
|
90
|
+
"""``True`` for compiler-generated temporaries."""
|
|
91
|
+
return self.name.startswith("%")
|
|
92
|
+
|
|
93
|
+
def __str__(self) -> str:
|
|
94
|
+
return self.name
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
Operand = Union[Const, Var]
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
# --------------------------------------------------------------------------
|
|
101
|
+
# Operator tables (shared by the interpreter and constant folding)
|
|
102
|
+
# --------------------------------------------------------------------------
|
|
103
|
+
|
|
104
|
+
BINARY_OPS: dict[str, Callable[[Any, Any], Any]] = {
|
|
105
|
+
"+": operator.add,
|
|
106
|
+
"-": operator.sub,
|
|
107
|
+
"*": operator.mul,
|
|
108
|
+
"/": operator.truediv,
|
|
109
|
+
"//": operator.floordiv,
|
|
110
|
+
"%": operator.mod,
|
|
111
|
+
"**": operator.pow,
|
|
112
|
+
"<<": operator.lshift,
|
|
113
|
+
">>": operator.rshift,
|
|
114
|
+
"&": operator.and_,
|
|
115
|
+
"|": operator.or_,
|
|
116
|
+
"^": operator.xor,
|
|
117
|
+
"==": operator.eq,
|
|
118
|
+
"!=": operator.ne,
|
|
119
|
+
"<": operator.lt,
|
|
120
|
+
"<=": operator.le,
|
|
121
|
+
">": operator.gt,
|
|
122
|
+
">=": operator.ge,
|
|
123
|
+
"is": operator.is_,
|
|
124
|
+
"is not": operator.is_not,
|
|
125
|
+
}
|
|
126
|
+
"""Binary operator spelling -> Python implementation."""
|
|
127
|
+
|
|
128
|
+
UNARY_OPS: dict[str, Callable[[Any], Any]] = {
|
|
129
|
+
"-": operator.neg,
|
|
130
|
+
"+": operator.pos,
|
|
131
|
+
"~": operator.invert,
|
|
132
|
+
"not": operator.not_,
|
|
133
|
+
}
|
|
134
|
+
"""Unary operator spelling -> Python implementation."""
|
|
135
|
+
|
|
136
|
+
BUILTINS: dict[str, Callable[..., Any]] = {
|
|
137
|
+
"abs": abs,
|
|
138
|
+
"min": min,
|
|
139
|
+
"max": max,
|
|
140
|
+
"int": int,
|
|
141
|
+
"float": float,
|
|
142
|
+
"bool": bool,
|
|
143
|
+
"str": str,
|
|
144
|
+
"round": round,
|
|
145
|
+
"print": print,
|
|
146
|
+
}
|
|
147
|
+
"""Python builtins callable from IR code (by name, via :class:`Call`)."""
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
# --------------------------------------------------------------------------
|
|
151
|
+
# Instructions
|
|
152
|
+
# --------------------------------------------------------------------------
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
@dataclass
|
|
156
|
+
class Copy:
|
|
157
|
+
"""``dest = src``"""
|
|
158
|
+
|
|
159
|
+
dest: Var
|
|
160
|
+
src: Operand
|
|
161
|
+
|
|
162
|
+
def uses(self) -> list[Operand]:
|
|
163
|
+
"""Operands read by this instruction."""
|
|
164
|
+
return [self.src]
|
|
165
|
+
|
|
166
|
+
def __str__(self) -> str:
|
|
167
|
+
return f"{self.dest} = {self.src}"
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
@dataclass
|
|
171
|
+
class BinOp:
|
|
172
|
+
"""``dest = left <op> right`` for any operator in :data:`BINARY_OPS`."""
|
|
173
|
+
|
|
174
|
+
dest: Var
|
|
175
|
+
op: str
|
|
176
|
+
left: Operand
|
|
177
|
+
right: Operand
|
|
178
|
+
|
|
179
|
+
def uses(self) -> list[Operand]:
|
|
180
|
+
"""Operands read by this instruction."""
|
|
181
|
+
return [self.left, self.right]
|
|
182
|
+
|
|
183
|
+
def __str__(self) -> str:
|
|
184
|
+
return f"{self.dest} = {self.left} {self.op} {self.right}"
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
@dataclass
|
|
188
|
+
class UnaryOp:
|
|
189
|
+
"""``dest = <op> operand`` for any operator in :data:`UNARY_OPS`."""
|
|
190
|
+
|
|
191
|
+
dest: Var
|
|
192
|
+
op: str
|
|
193
|
+
operand: Operand
|
|
194
|
+
|
|
195
|
+
def uses(self) -> list[Operand]:
|
|
196
|
+
"""Operands read by this instruction."""
|
|
197
|
+
return [self.operand]
|
|
198
|
+
|
|
199
|
+
def __str__(self) -> str:
|
|
200
|
+
sep = " " if self.op == "not" else ""
|
|
201
|
+
return f"{self.dest} = {self.op}{sep}{self.operand}"
|
|
202
|
+
|
|
203
|
+
|
|
204
|
+
@dataclass
|
|
205
|
+
class Call:
|
|
206
|
+
"""``dest = call func(args...)``.
|
|
207
|
+
|
|
208
|
+
``func`` names either another function of the same :class:`Module` or an
|
|
209
|
+
entry of :data:`BUILTINS` (module functions take precedence).
|
|
210
|
+
"""
|
|
211
|
+
|
|
212
|
+
dest: Var
|
|
213
|
+
func: str
|
|
214
|
+
args: list[Operand]
|
|
215
|
+
|
|
216
|
+
def uses(self) -> list[Operand]:
|
|
217
|
+
"""Operands read by this instruction."""
|
|
218
|
+
return list(self.args)
|
|
219
|
+
|
|
220
|
+
def __str__(self) -> str:
|
|
221
|
+
args = ", ".join(str(a) for a in self.args)
|
|
222
|
+
return f"{self.dest} = call {self.func}({args})"
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
Instruction = Union[Copy, BinOp, UnaryOp, Call]
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
@dataclass
|
|
229
|
+
class Jump:
|
|
230
|
+
"""Unconditional transfer of control to ``target``."""
|
|
231
|
+
|
|
232
|
+
target: str
|
|
233
|
+
|
|
234
|
+
def uses(self) -> list[Operand]:
|
|
235
|
+
"""Operands read by this terminator."""
|
|
236
|
+
return []
|
|
237
|
+
|
|
238
|
+
def successors(self) -> list[str]:
|
|
239
|
+
"""Labels this terminator may transfer control to."""
|
|
240
|
+
return [self.target]
|
|
241
|
+
|
|
242
|
+
def __str__(self) -> str:
|
|
243
|
+
return f"jump {self.target}"
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
@dataclass
|
|
247
|
+
class Branch:
|
|
248
|
+
"""Go to ``if_true`` when ``cond`` is truthy, otherwise to ``if_false``."""
|
|
249
|
+
|
|
250
|
+
cond: Operand
|
|
251
|
+
if_true: str
|
|
252
|
+
if_false: str
|
|
253
|
+
|
|
254
|
+
def uses(self) -> list[Operand]:
|
|
255
|
+
"""Operands read by this terminator."""
|
|
256
|
+
return [self.cond]
|
|
257
|
+
|
|
258
|
+
def successors(self) -> list[str]:
|
|
259
|
+
"""Labels this terminator may transfer control to."""
|
|
260
|
+
return [self.if_true, self.if_false]
|
|
261
|
+
|
|
262
|
+
def __str__(self) -> str:
|
|
263
|
+
return f"branch {self.cond}, {self.if_true}, {self.if_false}"
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
@dataclass
|
|
267
|
+
class Return:
|
|
268
|
+
"""Return ``value`` from the current function."""
|
|
269
|
+
|
|
270
|
+
value: Operand
|
|
271
|
+
|
|
272
|
+
def uses(self) -> list[Operand]:
|
|
273
|
+
"""Operands read by this terminator."""
|
|
274
|
+
return [self.value]
|
|
275
|
+
|
|
276
|
+
def successors(self) -> list[str]:
|
|
277
|
+
"""Labels this terminator may transfer control to (none)."""
|
|
278
|
+
return []
|
|
279
|
+
|
|
280
|
+
def __str__(self) -> str:
|
|
281
|
+
return f"return {self.value}"
|
|
282
|
+
|
|
283
|
+
|
|
284
|
+
Terminator = Union[Jump, Branch, Return]
|
|
285
|
+
|
|
286
|
+
|
|
287
|
+
# --------------------------------------------------------------------------
|
|
288
|
+
# Containers
|
|
289
|
+
# --------------------------------------------------------------------------
|
|
290
|
+
|
|
291
|
+
|
|
292
|
+
@dataclass
|
|
293
|
+
class BasicBlock:
|
|
294
|
+
"""A label, a list of instructions and a single terminator."""
|
|
295
|
+
|
|
296
|
+
label: str
|
|
297
|
+
instructions: list[Instruction] = field(default_factory=list)
|
|
298
|
+
terminator: Terminator | None = None
|
|
299
|
+
|
|
300
|
+
def successors(self) -> list[str]:
|
|
301
|
+
"""Labels of the blocks control may flow to next."""
|
|
302
|
+
return self.terminator.successors() if self.terminator is not None else []
|
|
303
|
+
|
|
304
|
+
|
|
305
|
+
@dataclass
|
|
306
|
+
class Function:
|
|
307
|
+
"""A function: parameter names plus an ordered label -> block mapping."""
|
|
308
|
+
|
|
309
|
+
name: str
|
|
310
|
+
params: list[str]
|
|
311
|
+
blocks: dict[str, BasicBlock] = field(default_factory=dict)
|
|
312
|
+
|
|
313
|
+
@property
|
|
314
|
+
def entry(self) -> BasicBlock:
|
|
315
|
+
"""The entry block (the first block inserted)."""
|
|
316
|
+
return next(iter(self.blocks.values()))
|
|
317
|
+
|
|
318
|
+
def __iter__(self) -> Iterator[BasicBlock]:
|
|
319
|
+
return iter(self.blocks.values())
|
|
320
|
+
|
|
321
|
+
def predecessors(self) -> dict[str, list[str]]:
|
|
322
|
+
"""Map every label to the labels of its predecessor blocks."""
|
|
323
|
+
preds: dict[str, list[str]] = {label: [] for label in self.blocks}
|
|
324
|
+
for block in self.blocks.values():
|
|
325
|
+
for succ in block.successors():
|
|
326
|
+
if succ in preds and block.label not in preds[succ]:
|
|
327
|
+
preds[succ].append(block.label)
|
|
328
|
+
return preds
|
|
329
|
+
|
|
330
|
+
def reachable(self) -> list[str]:
|
|
331
|
+
"""Labels reachable from the entry block, in reverse post-order."""
|
|
332
|
+
seen: set[str] = set()
|
|
333
|
+
order: list[str] = []
|
|
334
|
+
|
|
335
|
+
def visit(label: str) -> None:
|
|
336
|
+
# Iterative DFS to avoid recursion limits on large CFGs.
|
|
337
|
+
stack: list[tuple[str, Iterator[str]]] = [(label, iter(self.blocks[label].successors()))]
|
|
338
|
+
seen.add(label)
|
|
339
|
+
while stack:
|
|
340
|
+
current, succs = stack[-1]
|
|
341
|
+
for succ in succs:
|
|
342
|
+
if succ not in seen and succ in self.blocks:
|
|
343
|
+
seen.add(succ)
|
|
344
|
+
stack.append((succ, iter(self.blocks[succ].successors())))
|
|
345
|
+
break
|
|
346
|
+
else:
|
|
347
|
+
stack.pop()
|
|
348
|
+
order.append(current)
|
|
349
|
+
|
|
350
|
+
if self.blocks:
|
|
351
|
+
visit(self.entry.label)
|
|
352
|
+
order.reverse()
|
|
353
|
+
return order
|
|
354
|
+
|
|
355
|
+
def instruction_count(self) -> int:
|
|
356
|
+
"""Total number of instructions plus terminators across all blocks."""
|
|
357
|
+
return sum(len(b.instructions) + (b.terminator is not None) for b in self)
|
|
358
|
+
|
|
359
|
+
def verify(self) -> None:
|
|
360
|
+
"""Check structural invariants.
|
|
361
|
+
|
|
362
|
+
Raises:
|
|
363
|
+
IRVerificationError: if the function has no blocks, a block lacks a
|
|
364
|
+
terminator, a label is inconsistent, a jump targets an unknown
|
|
365
|
+
block, or an operator is unknown.
|
|
366
|
+
"""
|
|
367
|
+
if not self.blocks:
|
|
368
|
+
raise IRVerificationError(f"{self.name}: function has no blocks")
|
|
369
|
+
for label, block in self.blocks.items():
|
|
370
|
+
if block.label != label:
|
|
371
|
+
raise IRVerificationError(f"{self.name}: block key {label!r} != label {block.label!r}")
|
|
372
|
+
if block.terminator is None:
|
|
373
|
+
raise IRVerificationError(f"{self.name}: block {label!r} has no terminator")
|
|
374
|
+
for succ in block.successors():
|
|
375
|
+
if succ not in self.blocks:
|
|
376
|
+
raise IRVerificationError(f"{self.name}: block {label!r} jumps to unknown block {succ!r}")
|
|
377
|
+
for inst in block.instructions:
|
|
378
|
+
if isinstance(inst, BinOp) and inst.op not in BINARY_OPS:
|
|
379
|
+
raise IRVerificationError(f"{self.name}: unknown binary operator {inst.op!r}")
|
|
380
|
+
if isinstance(inst, UnaryOp) and inst.op not in UNARY_OPS:
|
|
381
|
+
raise IRVerificationError(f"{self.name}: unknown unary operator {inst.op!r}")
|
|
382
|
+
|
|
383
|
+
def __str__(self) -> str:
|
|
384
|
+
from .printer import format_function
|
|
385
|
+
|
|
386
|
+
return format_function(self)
|
|
387
|
+
|
|
388
|
+
|
|
389
|
+
@dataclass
|
|
390
|
+
class Module:
|
|
391
|
+
"""A collection of functions that may call each other by name."""
|
|
392
|
+
|
|
393
|
+
functions: dict[str, Function] = field(default_factory=dict)
|
|
394
|
+
|
|
395
|
+
def __getitem__(self, name: str) -> Function:
|
|
396
|
+
return self.functions[name]
|
|
397
|
+
|
|
398
|
+
def __contains__(self, name: object) -> bool:
|
|
399
|
+
return name in self.functions
|
|
400
|
+
|
|
401
|
+
def __iter__(self) -> Iterator[Function]:
|
|
402
|
+
return iter(self.functions.values())
|
|
403
|
+
|
|
404
|
+
def verify(self) -> None:
|
|
405
|
+
"""Verify every function (see :meth:`Function.verify`)."""
|
|
406
|
+
for fn in self:
|
|
407
|
+
fn.verify()
|
|
408
|
+
|
|
409
|
+
def __str__(self) -> str:
|
|
410
|
+
from .printer import format_module
|
|
411
|
+
|
|
412
|
+
return format_module(self)
|