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 pprint
import pstats import pstats
import sys import sys
from collections.abc import Mapping
from copy import deepcopy from copy import deepcopy
from dataclasses import asdict, dataclass, fields, is_dataclass from dataclasses import asdict, dataclass, fields, is_dataclass
from pstats import SortKey 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 import Lark, Token, Transformer, Tree, ast_utils, tree, v_args
from lark.tree import Meta from lark.tree import Meta
@@ -22,12 +24,138 @@ this_module = sys.modules[__name__]
pp = pprint.PrettyPrinter(indent=4, width=50) 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 # type Ident = str
# @dataclass # @dataclass
# class Ident(ast_utils.Ast): # class Ident(ast_utils.Ast):
# ident: str # ident: str
# def pretty_print(self): # def pretty_print(self, column : int = 0):
# if isinstance(self.ident, Token): # if isinstance(self.ident, Token):
# return f"{self.ident.value}" # return f"{self.ident.value}"
# else: # else:
@@ -38,23 +166,25 @@ pp = pprint.PrettyPrinter(indent=4, width=50)
class TypeDecl(ast_utils.Ast): class TypeDecl(ast_utils.Ast):
ident: Token ident: Token
def pretty_print(self): def pretty_print(self, column: int = 0):
return "type " + self.ident.value + "." template = "type {ident}."
return pretty_format(self, template, column=column)
@dataclass @dataclass
class Typeid(ast_utils.Ast): class Typeid(ast_utils.Ast):
ident: Token ident: Token
def pretty_print(self): def pretty_print(self, column: int = 0):
return self.ident.value template = "{ident}"
return pretty_format(self, template, column=column)
# @dataclass # @dataclass
# class Infix(ast_utils.Ast): # class Infix(ast_utils.Ast):
# infix: Token # infix: Token
# def pretty_print(self): # def pretty_print(self, column : int = 0):
# return self.infix.value # return self.infix.value
@@ -62,27 +192,9 @@ class Typeid(ast_utils.Ast):
class Pterm(ast_utils.Ast, ast_utils.AsList): class Pterm(ast_utils.Ast, ast_utils.AsList):
pterm: Ident | int | List pterm: Ident | int | List
def pretty_print(self): def pretty_print(self, column: int = 0):
if type(self.pterm) is int: template = "{pterm:pterm}"
result = str(self.pterm) return pretty_format(self, template, column=column)
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
@dataclass @dataclass
@@ -91,16 +203,15 @@ class Gbinding(ast_utils.Ast):
gterm: Gterm gterm: Gterm
gbinding: Optional[Gbinding] = None gbinding: Optional[Gbinding] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
result = "" result = ""
if type(self.value) is int: if type(self.value) is int:
result = "!" + str(self.value) + "=" + self.gterm.pretty_print() result = pretty_format(self, "!{value}={gterm}", column=column)
else: else:
result = str(self.value) + "=" + self.gterm.pretty_print() result = pretty_format(self, "{value}={gterm}", column=column)
if self.gbinding is not None: if self.gbinding is not None:
result += ";" result += pretty_format(self, ";{gbinding}", column=column)
result += self.gbinding.pretty_print()
return result return result
@@ -109,18 +220,19 @@ class Gbinding(ast_utils.Ast):
class IdentGterm(ast_utils.Ast): class IdentGterm(ast_utils.Ast):
ident_gterm: Ident ident_gterm: Ident
def pretty_print(self): def pretty_print(self, column: int = 0):
return self.ident_gterm.value return pretty_format(self, "{ident_gterm}", column=column)
@dataclass @dataclass
class GtermList(ast_utils.Ast, ast_utils.AsList): class GtermList(ast_utils.Ast, ast_utils.AsList):
gterms: Optional[List[Gterm]] = None 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: if self.gterms is None:
return "" return ""
return ", ".join([t.pretty_print() for t in self.gterms]) return pretty_format(self, "{gterms:gterm}", column=column)
@dataclass @dataclass
@@ -133,7 +245,7 @@ class FunGterm(ast_utils.Ast):
phase: Optional[int] = None phase: Optional[int] = None
at: Optional[Ident] = None at: Optional[Ident] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
str = "" str = ""
str += self.fun_gterm.value str += self.fun_gterm.value
str += "(" str += "("
@@ -152,7 +264,7 @@ class InfixGterm(ast_utils.Ast):
infix: Infix infix: Infix
second_infix_gterm: Gterm 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()}" 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): class ChoiceGterm(ast_utils.Ast):
choice_gterm: Tuple[Gterm, Gterm] choice_gterm: Tuple[Gterm, Gterm]
def pretty_print(self): def pretty_print(self, column: int = 0):
left, right = self.choice_gterm left, right = self.choice_gterm
return f"choice [ {left.pretty_print()}, {right.pretty_print()} ]" return f"choice [ {left.pretty_print()}, {right.pretty_print()} ]"
@@ -171,7 +283,7 @@ class ArithGterm(ast_utils.Ast):
operand: str operand: str
value: int value: int
def pretty_print(self): def pretty_print(self, column: int = 0):
return f"{self.arith_gterm.pretty_print()} {self.operand} {self.value}" return f"{self.arith_gterm.pretty_print()} {self.operand} {self.value}"
@@ -180,7 +292,7 @@ class Arith2Gterm(ast_utils.Ast):
value: int value: int
arith_gterm: Gterm arith_gterm: Gterm
def pretty_print(self): def pretty_print(self, column: int = 0):
return f"{self.value} + {self.arith_gterm.pretty_print()}" return f"{self.value} + {self.arith_gterm.pretty_print()}"
@@ -189,7 +301,7 @@ class InjeventGterm(ast_utils.Ast):
event_gterms: GtermList event_gterms: GtermList
at: Optional[Ident] = None at: Optional[Ident] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
result = f"inj-event ( {self.event_gterms.pretty_print()} )" result = f"inj-event ( {self.event_gterms.pretty_print()} )"
if self.at is not None: if self.at is not None:
result += f"@ {self.at}" result += f"@ {self.at}"
@@ -201,7 +313,7 @@ class ImpliesGterm(ast_utils.Ast):
left: Gterm left: Gterm
right: Gterm right: Gterm
def pretty_print(self): def pretty_print(self, column: int = 0):
return f"{self.left.pretty_print()} ==> {self.right.pretty_print()}" return f"{self.left.pretty_print()} ==> {self.right.pretty_print()}"
@@ -210,7 +322,7 @@ class EventGterm(ast_utils.Ast):
event_gterms: GtermList event_gterms: GtermList
at: Optional[Ident] = None at: Optional[Ident] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
str = "event (" str = "event ("
str += self.event_gterms.pretty_print() str += self.event_gterms.pretty_print()
str += ")" str += ")"
@@ -223,7 +335,7 @@ class EventGterm(ast_utils.Ast):
class ParenGterm(ast_utils.Ast): class ParenGterm(ast_utils.Ast):
paren_gterms: GtermList paren_gterms: GtermList
def pretty_print(self): def pretty_print(self, column: int = 0):
return "(" + self.paren_gterms.pretty_print() + ")" return "(" + self.paren_gterms.pretty_print() + ")"
@@ -233,7 +345,7 @@ class LetGterm(ast_utils.Ast):
first_gterm: Gterm first_gterm: Gterm
second_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()}" 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 ident: Ident
gbinding: Optional[Gbinding] = None gbinding: Optional[Gbinding] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
str = f"new {self.ident.value}" str = f"new {self.ident.value}"
if self.gbinding is not None: if self.gbinding is not None:
str += "[" str += "["
@@ -268,7 +380,7 @@ class Gterm(ast_utils.Ast):
| LetGterm | LetGterm
) )
def pretty_print(self): def pretty_print(self, column: int = 0):
if isinstance( if isinstance(
self.gterm, self.gterm,
( (
@@ -298,7 +410,7 @@ class Typedecl(ast_utils.Ast):
optional_typedecl: Optional[Typedecl] = None optional_typedecl: Optional[Typedecl] = None
# _non_empty_seq{IDENT} ":" typeid [ "," typedecl ] # _non_empty_seq{IDENT} ":" typeid [ "," typedecl ]
def pretty_print(self): def pretty_print(self, column: int = 0):
str = "" str = ""
# str += ", ".join([t.pretty_print() for t in self.type_list]) # str += ", ".join([t.pretty_print() for t in self.type_list])
if isinstance(self.type_list, List): if isinstance(self.type_list, List):
@@ -319,7 +431,7 @@ class LetfunDecl(ast_utils.Ast):
typedecl: Optional[Typedecl] typedecl: Optional[Typedecl]
pterm: Pterm pterm: Pterm
def pretty_print(self): def pretty_print(self, column: int = 0):
str = f"letfun {self.ident.value}" str = f"letfun {self.ident.value}"
if self.typedecl is not None: if self.typedecl is not None:
str += "(" str += "("
@@ -342,7 +454,7 @@ class LemmaGterm(ast_utils.Ast):
gterm: Gterm gterm: Gterm
lemma: Optional[Lemma] = None lemma: Optional[Lemma] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
str = self.gterm.pretty_print() str = self.gterm.pretty_print()
if self.lemma is not None: if self.lemma is not None:
str += ";" + self.lemma.pretty_print() str += ";" + self.lemma.pretty_print()
@@ -355,7 +467,7 @@ class LemmaPublicVars(ast_utils.Ast):
public_vars: List public_vars: List
lemma: Optional[Lemma] = None lemma: Optional[Lemma] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
str = self.gterm.pretty_print() str = self.gterm.pretty_print()
str += "for { public_vars" str += "for { public_vars"
str += ",".join( str += ",".join(
@@ -377,7 +489,7 @@ class LemmaPublicVarsSecret(ast_utils.Ast):
public_vars: List public_vars: List
lemma: Optional[Lemma] = None lemma: Optional[Lemma] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
str = self.gterm.pretty_print() str = self.gterm.pretty_print()
str += "for { secret" str += "for { secret"
str += self.secret.value str += self.secret.value
@@ -397,7 +509,7 @@ class LemmaPublicVarsSecret(ast_utils.Ast):
class Lemma(ast_utils.Ast): class Lemma(ast_utils.Ast):
lemma: LemmaGterm | LemmaPublicVars | LemmaPublicVarsSecret lemma: LemmaGterm | LemmaPublicVars | LemmaPublicVarsSecret
def pretty_print(self): def pretty_print(self, column: int = 0):
return self.lemma.pretty_print() return self.lemma.pretty_print()
# def __post_init__(self): # def __post_init__(self):
@@ -416,7 +528,7 @@ class LemmaDeclCore(ast_utils.Ast):
lemma: Lemma lemma: Lemma
# lemma_decl_core: "lemma" [ typedecl ";"] lemma "." # lemma_decl_core: "lemma" [ typedecl ";"] lemma "."
def pretty_print(self): def pretty_print(self, column: int = 0):
str = "" str = ""
str += "lemma " str += "lemma "
if self.typedecl is not None: if self.typedecl is not None:
@@ -431,7 +543,7 @@ class LemmaDeclCore(ast_utils.Ast):
class LemmaAnnotation(ast_utils.Ast): class LemmaAnnotation(ast_utils.Ast):
annotation: str annotation: str
def pretty_print(self): def pretty_print(self, column: int = 0):
str = "" str = ""
str += self.annotation str += self.annotation
return str return str
@@ -442,7 +554,7 @@ class LemmaDecl(ast_utils.Ast):
lemma_decl_annotation: Optional[LemmaAnnotation] lemma_decl_annotation: Optional[LemmaAnnotation]
lemma_decl_core: LemmaDeclCore lemma_decl_core: LemmaDeclCore
def pretty_print(self): def pretty_print(self, column: int = 0):
str = "" str = ""
if self.lemma_decl_annotation is not None: if self.lemma_decl_annotation is not None:
str += f"@lemma {self.lemma_decl_annotation}\n" str += f"@lemma {self.lemma_decl_annotation}\n"
@@ -464,7 +576,7 @@ class LemmaDecl(ast_utils.Ast):
class Query(ast_utils.Ast): class Query(ast_utils.Ast):
query: QueryGterm | QuerySecret | QueryPutBegin query: QueryGterm | QuerySecret | QueryPutBegin
def pretty_print(self): def pretty_print(self, column: int = 0):
return self.query.pretty_print() return self.query.pretty_print()
@@ -472,7 +584,7 @@ class Query(ast_utils.Ast):
class QueryAnnotation(ast_utils.Ast): class QueryAnnotation(ast_utils.Ast):
annotation: str annotation: str
def pretty_print(self): def pretty_print(self, column: int = 0):
return self.annotation return self.annotation
@@ -480,7 +592,7 @@ class QueryAnnotation(ast_utils.Ast):
class ReachableAnnotation(ast_utils.Ast): class ReachableAnnotation(ast_utils.Ast):
annotation: str annotation: str
def pretty_print(self): def pretty_print(self, column: int = 0):
return self.annotation return self.annotation
@@ -489,7 +601,7 @@ class QueryDeclCore(ast_utils.Ast):
typedecl: Optional[Typedecl] typedecl: Optional[Typedecl]
query: Query query: Query
def pretty_print(self): def pretty_print(self, column: int = 0):
str = "query " str = "query "
if self.typedecl is not None: if self.typedecl is not None:
str += self.typedecl.pretty_print() str += self.typedecl.pretty_print()
@@ -504,7 +616,7 @@ class QueryDecl(ast_utils.Ast):
query_decl_annotation: Optional[QueryAnnotation | ReachableAnnotation] query_decl_annotation: Optional[QueryAnnotation | ReachableAnnotation]
query_decl_core: QueryDeclCore query_decl_core: QueryDeclCore
def pretty_print(self): def pretty_print(self, column: int = 0):
str = "" str = ""
if self.query_decl_annotation is not None: if self.query_decl_annotation is not None:
if isinstance(self.query_decl_annotation, QueryAnnotation): if isinstance(self.query_decl_annotation, QueryAnnotation):
@@ -521,7 +633,7 @@ class QueryGterm(ast_utils.Ast):
public_vars: Optional[List[Ident]] = None public_vars: Optional[List[Ident]] = None
query: Optional[Query] = None query: Optional[Query] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
str = self.gterm.pretty_print() str = self.gterm.pretty_print()
if self.public_vars is not None: if self.public_vars is not None:
str += "public_vars " str += "public_vars "
@@ -542,7 +654,7 @@ class QuerySecret(ast_utils.Ast):
public_vars: Optional[List[Ident]] = None public_vars: Optional[List[Ident]] = None
query: Optional[Query] = None query: Optional[Query] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
str = "secret" str = "secret"
str += self.ident.pretty_print() str += self.ident.pretty_print()
if self.public_vars is not None: if self.public_vars is not None:
@@ -563,7 +675,7 @@ class QueryPutBegin(ast_utils.Ast):
event_list: List[Ident] event_list: List[Ident]
query: Optional[Query] = None query: Optional[Query] = None
def pretty_print(self): def pretty_print(self, column: int = 0):
str = "putbegin event :" str = "putbegin event :"
str += ", ".join( str += ", ".join(
[ [
@@ -580,7 +692,7 @@ class QueryPutBegin(ast_utils.Ast):
class Decl(ast_utils.Ast): class Decl(ast_utils.Ast):
decl: LemmaDecl | QueryDecl | TypeDecl | LetfunDecl decl: LemmaDecl | QueryDecl | TypeDecl | LetfunDecl
def pretty_print(self): def pretty_print(self, column: int = 0):
return self.decl.pretty_print() return self.decl.pretty_print()