mirror of
https://github.com/rosenpass/rosenpass.git
synced 2026-07-30 23:30:11 -07:00
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:
co-authored by
Anja Rabich
parent
a7944a5935
commit
a6cf1c500c
+179
-67
@@ -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()
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user