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 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)