feat(WIP): better templating for pretty printer

with configuration options

Co-authored-by: Anja Rabich <a.rabich@uni-luebeck.de>
This commit is contained in:
Benjamin Lipp
2026-06-02 18:02:14 +02:00
co-authored by Anja Rabich
parent a7944a5935
commit a6cf1c500c
+179 -67
View File
@@ -6,10 +6,12 @@ import io
import pprint
import pstats
import sys
from collections.abc import Mapping
from copy import deepcopy
from dataclasses import asdict, dataclass, fields, is_dataclass
from pstats import SortKey
from typing import List, Optional, Tuple
from string import Formatter
from typing import Any, List, Optional, Tuple
from lark import Lark, Token, Transformer, Tree, ast_utils, tree, v_args
from lark.tree import Meta
@@ -22,12 +24,138 @@ this_module = sys.modules[__name__]
pp = pprint.PrettyPrinter(indent=4, width=50)
CONFIG = {
"space": " ",
"pterm": {
"list_separator": ", ",
"left_bracket": "(",
"right_bracket": ")",
"empty_brackets": True,
},
"gterm": {
"list_separator": ", ",
"left_bracket": "",
"right_bracket": "",
"empty_brackets": False,
},
"infix_separator": " ",
"line_break": "\n",
}
class AttrMap:
"""
Small wrapper so that dict values can be accessed as {ctx.foo}
instead of only {ctx[foo]}.
"""
def __init__(self, mapping: Mapping[str, Any]):
self._mapping = mapping
def __getattr__(self, name: str) -> Any:
try:
return self._mapping[name]
except KeyError:
raise AttributeError(name) from None
def __getitem__(self, name: str) -> Any:
return self._mapping[name]
def __contains__(self, name: str) -> bool:
return name in self._mapping
def __str__(self) -> str:
return str(self._mapping)
def get_list_config(format_spec: str, ctx: Mapping[str, Any] | None = None):
if ctx:
if format_spec in ctx:
list_config = ctx[format_spec]
return (
list_config["left_bracket"],
list_config["right_bracket"],
list_config["list_separator"],
list_config["empty_brackets"],
)
return None
def pretty(
value: Any, column: int, format_spec: str, ctx: Mapping[str, Any] | None = None
) -> str:
if isinstance(value, List):
left_bracket, right_bracket, list_separator, empty_brackets = get_list_config(
format_spec, ctx
)
if len(value) > 1:
return (
left_bracket
+ list_separator.join(
pretty(item, column=column, format_spec=format_spec, ctx=ctx)
for item in value
)
+ right_bracket
)
elif len(value) == 1:
return pretty(value[0], column=column, format_spec=format_spec, ctx=ctx)
else:
return left_bracket + right_bracket if empty_brackets else ""
pretty_print = getattr(value, "pretty_print", None)
if callable(pretty_print):
return pretty_print(column=column)
return str(value)
class PrettyFormatter(Formatter):
def __init__(
self,
root: Any,
*,
column: int,
ctx: Mapping[str, Any] | None = None,
):
super().__init__()
self.root = root
self.column = column
self.ctx = AttrMap(ctx or {})
def get_value(
self,
key: Any,
args: tuple[Any, ...],
kwargs: dict[str, Any],
) -> Any:
if key == "self":
return self.root
if key == "ctx":
return self.ctx
if isinstance(key, str):
return getattr(self.root, key)
return super().get_value(key, args, kwargs)
def format_field(self, value: Any, format_spec: str) -> str:
child_column = self.column
return pretty(value, column=child_column, format_spec=format_spec, ctx=self.ctx)
def pretty_format(obj: Any, template: str, *, column: int = 0) -> str:
return PrettyFormatter(obj, column=column, ctx=CONFIG).format(template)
# type Ident = str
# @dataclass
# class Ident(ast_utils.Ast):
# ident: str
# def pretty_print(self):
# def pretty_print(self, column : int = 0):
# if isinstance(self.ident, Token):
# return f"{self.ident.value}"
# else:
@@ -38,23 +166,25 @@ pp = pprint.PrettyPrinter(indent=4, width=50)
class TypeDecl(ast_utils.Ast):
ident: Token
def pretty_print(self):
return "type " + self.ident.value + "."
def pretty_print(self, column: int = 0):
template = "type {ident}."
return pretty_format(self, template, column=column)
@dataclass
class Typeid(ast_utils.Ast):
ident: Token
def pretty_print(self):
return self.ident.value
def pretty_print(self, column: int = 0):
template = "{ident}"
return pretty_format(self, template, column=column)
# @dataclass
# class Infix(ast_utils.Ast):
# infix: Token
# def pretty_print(self):
# def pretty_print(self, column : int = 0):
# return self.infix.value
@@ -62,27 +192,9 @@ class Typeid(ast_utils.Ast):
class Pterm(ast_utils.Ast, ast_utils.AsList):
pterm: Ident | int | List
def pretty_print(self):
if type(self.pterm) is int:
result = str(self.pterm)
elif isinstance(self.pterm, List):
# TODO: also add a pterm_list rule and only call pretty_print()?
result = (
"("
+ ",".join(
[
x.pretty_print() if hasattr(x, "pretty_print") else str(x)
for x in self.pterm
]
)
+ ")"
if len(self.pterm) > 1
else str(self.pterm[0])
)
else:
result = self.pterm.value
return result
def pretty_print(self, column: int = 0):
template = "{pterm:pterm}"
return pretty_format(self, template, column=column)
@dataclass
@@ -91,16 +203,15 @@ class Gbinding(ast_utils.Ast):
gterm: Gterm
gbinding: Optional[Gbinding] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
result = ""
if type(self.value) is int:
result = "!" + str(self.value) + "=" + self.gterm.pretty_print()
result = pretty_format(self, "!{value}={gterm}", column=column)
else:
result = str(self.value) + "=" + self.gterm.pretty_print()
result = pretty_format(self, "{value}={gterm}", column=column)
if self.gbinding is not None:
result += ";"
result += self.gbinding.pretty_print()
result += pretty_format(self, ";{gbinding}", column=column)
return result
@@ -109,18 +220,19 @@ class Gbinding(ast_utils.Ast):
class IdentGterm(ast_utils.Ast):
ident_gterm: Ident
def pretty_print(self):
return self.ident_gterm.value
def pretty_print(self, column: int = 0):
return pretty_format(self, "{ident_gterm}", column=column)
@dataclass
class GtermList(ast_utils.Ast, ast_utils.AsList):
gterms: Optional[List[Gterm]] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
# TODO: move the None case into the Formatter? none_repr config entry, and if clause in the beginning of pretty function
if self.gterms is None:
return ""
return ", ".join([t.pretty_print() for t in self.gterms])
return pretty_format(self, "{gterms:gterm}", column=column)
@dataclass
@@ -133,7 +245,7 @@ class FunGterm(ast_utils.Ast):
phase: Optional[int] = None
at: Optional[Ident] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = ""
str += self.fun_gterm.value
str += "("
@@ -152,7 +264,7 @@ class InfixGterm(ast_utils.Ast):
infix: Infix
second_infix_gterm: Gterm
def pretty_print(self):
def pretty_print(self, column: int = 0):
return f"{self.first_infix_gterm.pretty_print()} {self.infix} {self.second_infix_gterm.pretty_print()}"
@@ -160,7 +272,7 @@ class InfixGterm(ast_utils.Ast):
class ChoiceGterm(ast_utils.Ast):
choice_gterm: Tuple[Gterm, Gterm]
def pretty_print(self):
def pretty_print(self, column: int = 0):
left, right = self.choice_gterm
return f"choice [ {left.pretty_print()}, {right.pretty_print()} ]"
@@ -171,7 +283,7 @@ class ArithGterm(ast_utils.Ast):
operand: str
value: int
def pretty_print(self):
def pretty_print(self, column: int = 0):
return f"{self.arith_gterm.pretty_print()} {self.operand} {self.value}"
@@ -180,7 +292,7 @@ class Arith2Gterm(ast_utils.Ast):
value: int
arith_gterm: Gterm
def pretty_print(self):
def pretty_print(self, column: int = 0):
return f"{self.value} + {self.arith_gterm.pretty_print()}"
@@ -189,7 +301,7 @@ class InjeventGterm(ast_utils.Ast):
event_gterms: GtermList
at: Optional[Ident] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
result = f"inj-event ( {self.event_gterms.pretty_print()} )"
if self.at is not None:
result += f"@ {self.at}"
@@ -201,7 +313,7 @@ class ImpliesGterm(ast_utils.Ast):
left: Gterm
right: Gterm
def pretty_print(self):
def pretty_print(self, column: int = 0):
return f"{self.left.pretty_print()} ==> {self.right.pretty_print()}"
@@ -210,7 +322,7 @@ class EventGterm(ast_utils.Ast):
event_gterms: GtermList
at: Optional[Ident] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = "event ("
str += self.event_gterms.pretty_print()
str += ")"
@@ -223,7 +335,7 @@ class EventGterm(ast_utils.Ast):
class ParenGterm(ast_utils.Ast):
paren_gterms: GtermList
def pretty_print(self):
def pretty_print(self, column: int = 0):
return "(" + self.paren_gterms.pretty_print() + ")"
@@ -233,7 +345,7 @@ class LetGterm(ast_utils.Ast):
first_gterm: Gterm
second_gterm: Gterm
def pretty_print(self):
def pretty_print(self, column: int = 0):
return f"let {self.ident.value} = {self.first_gterm.pretty_print()} in\n {self.second_gterm.pretty_print()}"
@@ -242,7 +354,7 @@ class SampleGterm(ast_utils.Ast):
ident: Ident
gbinding: Optional[Gbinding] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = f"new {self.ident.value}"
if self.gbinding is not None:
str += "["
@@ -268,7 +380,7 @@ class Gterm(ast_utils.Ast):
| LetGterm
)
def pretty_print(self):
def pretty_print(self, column: int = 0):
if isinstance(
self.gterm,
(
@@ -298,7 +410,7 @@ class Typedecl(ast_utils.Ast):
optional_typedecl: Optional[Typedecl] = None
# _non_empty_seq{IDENT} ":" typeid [ "," typedecl ]
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = ""
# str += ", ".join([t.pretty_print() for t in self.type_list])
if isinstance(self.type_list, List):
@@ -319,7 +431,7 @@ class LetfunDecl(ast_utils.Ast):
typedecl: Optional[Typedecl]
pterm: Pterm
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = f"letfun {self.ident.value}"
if self.typedecl is not None:
str += "("
@@ -342,7 +454,7 @@ class LemmaGterm(ast_utils.Ast):
gterm: Gterm
lemma: Optional[Lemma] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = self.gterm.pretty_print()
if self.lemma is not None:
str += ";" + self.lemma.pretty_print()
@@ -355,7 +467,7 @@ class LemmaPublicVars(ast_utils.Ast):
public_vars: List
lemma: Optional[Lemma] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = self.gterm.pretty_print()
str += "for { public_vars"
str += ",".join(
@@ -377,7 +489,7 @@ class LemmaPublicVarsSecret(ast_utils.Ast):
public_vars: List
lemma: Optional[Lemma] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = self.gterm.pretty_print()
str += "for { secret"
str += self.secret.value
@@ -397,7 +509,7 @@ class LemmaPublicVarsSecret(ast_utils.Ast):
class Lemma(ast_utils.Ast):
lemma: LemmaGterm | LemmaPublicVars | LemmaPublicVarsSecret
def pretty_print(self):
def pretty_print(self, column: int = 0):
return self.lemma.pretty_print()
# def __post_init__(self):
@@ -416,7 +528,7 @@ class LemmaDeclCore(ast_utils.Ast):
lemma: Lemma
# lemma_decl_core: "lemma" [ typedecl ";"] lemma "."
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = ""
str += "lemma "
if self.typedecl is not None:
@@ -431,7 +543,7 @@ class LemmaDeclCore(ast_utils.Ast):
class LemmaAnnotation(ast_utils.Ast):
annotation: str
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = ""
str += self.annotation
return str
@@ -442,7 +554,7 @@ class LemmaDecl(ast_utils.Ast):
lemma_decl_annotation: Optional[LemmaAnnotation]
lemma_decl_core: LemmaDeclCore
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = ""
if self.lemma_decl_annotation is not None:
str += f"@lemma {self.lemma_decl_annotation}\n"
@@ -464,7 +576,7 @@ class LemmaDecl(ast_utils.Ast):
class Query(ast_utils.Ast):
query: QueryGterm | QuerySecret | QueryPutBegin
def pretty_print(self):
def pretty_print(self, column: int = 0):
return self.query.pretty_print()
@@ -472,7 +584,7 @@ class Query(ast_utils.Ast):
class QueryAnnotation(ast_utils.Ast):
annotation: str
def pretty_print(self):
def pretty_print(self, column: int = 0):
return self.annotation
@@ -480,7 +592,7 @@ class QueryAnnotation(ast_utils.Ast):
class ReachableAnnotation(ast_utils.Ast):
annotation: str
def pretty_print(self):
def pretty_print(self, column: int = 0):
return self.annotation
@@ -489,7 +601,7 @@ class QueryDeclCore(ast_utils.Ast):
typedecl: Optional[Typedecl]
query: Query
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = "query "
if self.typedecl is not None:
str += self.typedecl.pretty_print()
@@ -504,7 +616,7 @@ class QueryDecl(ast_utils.Ast):
query_decl_annotation: Optional[QueryAnnotation | ReachableAnnotation]
query_decl_core: QueryDeclCore
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = ""
if self.query_decl_annotation is not None:
if isinstance(self.query_decl_annotation, QueryAnnotation):
@@ -521,7 +633,7 @@ class QueryGterm(ast_utils.Ast):
public_vars: Optional[List[Ident]] = None
query: Optional[Query] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = self.gterm.pretty_print()
if self.public_vars is not None:
str += "public_vars "
@@ -542,7 +654,7 @@ class QuerySecret(ast_utils.Ast):
public_vars: Optional[List[Ident]] = None
query: Optional[Query] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = "secret"
str += self.ident.pretty_print()
if self.public_vars is not None:
@@ -563,7 +675,7 @@ class QueryPutBegin(ast_utils.Ast):
event_list: List[Ident]
query: Optional[Query] = None
def pretty_print(self):
def pretty_print(self, column: int = 0):
str = "putbegin event :"
str += ", ".join(
[
@@ -580,7 +692,7 @@ class QueryPutBegin(ast_utils.Ast):
class Decl(ast_utils.Ast):
decl: LemmaDecl | QueryDecl | TypeDecl | LetfunDecl
def pretty_print(self):
def pretty_print(self, column: int = 0):
return self.decl.pretty_print()