feat: pv pretty_printer for minimal grammar

Co-authored-by: Benjamin Lipp <blipp@rosenpass.eu>
This commit is contained in:
Anja Rabich
2026-05-26 17:57:18 +02:00
co-authored by Benjamin Lipp
parent 0f885a0c21
commit 36d3c8f6e3
+259 -62
View File
@@ -25,13 +25,13 @@ pp = pprint.PrettyPrinter(indent=4, width=50)
# 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):
# if isinstance(self.ident, Token): # if isinstance(self.ident, Token):
# return f"{self.ident.value}" # return f"{self.ident.value}"
# else: # else:
# return "whaaaaat" # return "whaaaaat"
@dataclass @dataclass
@@ -50,22 +50,59 @@ class Typeid(ast_utils.Ast):
return self.ident.value return self.ident.value
@dataclass # @dataclass
class Infix(ast_utils.Ast): # class Infix(ast_utils.Ast):
infix: Token # infix: Token
def pretty_print(self): # def pretty_print(self):
return self.infix.value # return self.infix.value
@dataclass @dataclass
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):
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
@dataclass @dataclass
class Gbinding(ast_utils.Ast): class Gbinding(ast_utils.Ast):
value: int | Ident value: int | Ident
gterm: Gterm
gbinding: Optional[Gbinding] = None
def pretty_print(self):
result = ""
if type(self.value) is int:
result = "!" + str(self.value) + "=" + self.gterm.pretty_print()
else:
result = str(self.value) + "=" + self.gterm.pretty_print()
if self.gbinding is not None:
result += ";"
result += self.gbinding.pretty_print()
return result
@dataclass @dataclass
@@ -73,15 +110,17 @@ class IdentGterm(ast_utils.Ast):
ident_gterm: Ident ident_gterm: Ident
def pretty_print(self): def pretty_print(self):
return self.ident_gterm.pretty_print() return self.ident_gterm.value
@dataclass @dataclass
class GtermList(ast_utils.Ast, ast_utils.AsList): class GtermList(ast_utils.Ast, ast_utils.AsList):
gterms: List[Gterm] gterms: Optional[List[Gterm]] = None
def pretty_print(self): def pretty_print(self):
return ",".join([t.pretty_print() for t in self.gterms]) if self.gterms is None:
return ""
return ", ".join([t.pretty_print() for t in self.gterms])
@dataclass @dataclass
@@ -96,14 +135,14 @@ class FunGterm(ast_utils.Ast):
def pretty_print(self): def pretty_print(self):
str = "" str = ""
str += self.fun_gterm.pretty_print() str += self.fun_gterm.value
str += "(" str += "("
str += self.gterm_list.pretty_print() str += self.gterm_list.pretty_print()
str += ")" str += ")"
if self.phase is not None: if self.phase is not None:
str += "not implemented" str += f"phase {self.phase}"
if self.at is not None: if self.at is not None:
str += "not implemented" str += f"@ {self.at}"
return str return str
@@ -113,32 +152,48 @@ class InfixGterm(ast_utils.Ast):
infix: Infix infix: Infix
second_infix_gterm: Gterm second_infix_gterm: Gterm
def pretty_print(self):
return f"{self.first_infix_gterm.pretty_print()} {self.infix} {self.second_infix_gterm.pretty_print()}"
@dataclass @dataclass
class ChoiceGterm(ast_utils.Ast): class ChoiceGterm(ast_utils.Ast):
choice_gterm: Optional[Tuple[Gterm, Gterm]] = None choice_gterm: Tuple[Gterm, Gterm]
def pretty_print(self):
left, right = self.choice_gterm
return f"choice [ {left.pretty_print()}, {right.pretty_print()} ]"
@dataclass @dataclass
class ArithGterm(ast_utils.Ast): class ArithGterm(ast_utils.Ast):
arith_gterm: Gterm arith_gterm: Gterm
operand: str operand: str
value: int | Gterm value: int
def pretty_print(self):
return f"{self.arith_gterm.pretty_print()} {self.operand} {self.value}"
@dataclass @dataclass
class Arith2Gterm(ast_utils.Ast): class Arith2Gterm(ast_utils.Ast):
value: int | Gterm value: int
arith_gterm: Gterm arith_gterm: Gterm
def pretty_print(self):
return f"{self.value} + {self.arith_gterm.pretty_print()}"
@dataclass @dataclass
class InjEventGterm(ast_utils.Ast): 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):
return "" result = f"inj-event ( {self.event_gterms.pretty_print()} )"
if self.at is not None:
result += f"@ {self.at}"
return result
@dataclass @dataclass
@@ -146,8 +201,8 @@ class ImpliesGterm(ast_utils.Ast):
left: Gterm left: Gterm
right: Gterm right: Gterm
def pretty_print(self, column=0, indent=2): def pretty_print(self):
return f"{self.left.pretty_print(column, indent)} ==> {self.right.pretty_print(column, indent)}" return f"{self.left.pretty_print()} ==> {self.right.pretty_print()}"
@dataclass @dataclass
@@ -155,10 +210,21 @@ class EventGterm(ast_utils.Ast):
event_gterms: GtermList event_gterms: GtermList
at: Optional[Ident] = None at: Optional[Ident] = None
def pretty_print(self):
str = "event ("
str += self.event_gterms.pretty_print()
str += ")"
if self.at is not None:
str += "@" + self.at.value
return str
@dataclass @dataclass
class ParenGterm(ast_utils.Ast): class ParenGterm(ast_utils.Ast):
paren_gterms: Optional[GtermList] = None paren_gterms: GtermList
def pretty_print(self):
return "(" + self.paren_gterms.pretty_print() + ")"
@dataclass @dataclass
@@ -167,11 +233,22 @@ class LetGterm(ast_utils.Ast):
first_gterm: Gterm first_gterm: Gterm
second_gterm: Gterm second_gterm: Gterm
def pretty_print(self):
return f"let {self.ident.value} = {self.first_gterm.pretty_print()} in\n {self.second_gterm.pretty_print()}"
@dataclass @dataclass
class SampleGterm(ast_utils.Ast): class SampleGterm(ast_utils.Ast):
ident: Ident ident: Ident
gbintion: Optional[Gbinding] = None gbinding: Optional[Gbinding] = None
def pretty_print(self):
str = f"new {self.ident.value}"
if self.gbinding is not None:
str += "["
str += self.gbinding.pretty_print()
str += "]"
return str
@dataclass @dataclass
@@ -183,7 +260,7 @@ class Gterm(ast_utils.Ast):
| ChoiceGterm | ChoiceGterm
| ArithGterm | ArithGterm
| Arith2Gterm | Arith2Gterm
| InjEventGterm | InjeventGterm
| ImpliesGterm | ImpliesGterm
| EventGterm | EventGterm
| ParenGterm | ParenGterm
@@ -192,7 +269,23 @@ class Gterm(ast_utils.Ast):
) )
def pretty_print(self): def pretty_print(self):
if isinstance(self.gterm, (IdentGterm, FunGterm)): if isinstance(
self.gterm,
(
IdentGterm,
FunGterm,
InfixGterm,
ChoiceGterm,
ArithGterm,
Arith2Gterm,
InjeventGterm,
ImpliesGterm,
EventGterm,
ParenGterm,
SampleGterm,
LetGterm,
),
):
return self.gterm.pretty_print() return self.gterm.pretty_print()
return "not implemented" return "not implemented"
# return f"{' ' * indent}{self}" # return f"{' ' * indent}{self}"
@@ -202,26 +295,47 @@ class Gterm(ast_utils.Ast):
class Typedecl(ast_utils.Ast): class Typedecl(ast_utils.Ast):
type_list: List[Ident] type_list: List[Ident]
typeid: Typeid typeid: Typeid
optional_typedecl: Optional[Type] = 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):
str = "" str = ""
# str += ", ".join([t.pretty_print() for t in self.type_list]) # str += ", ".join([t.pretty_print() for t in self.type_list])
str += ", ".join([t for t in self.type_list]) if isinstance(self.type_list, List):
str += ":" str += ", ".join([t for t in self.type_list])
elif isinstance(self.type_list, Token):
str += self.type_list.value
str += ": "
str += self.typeid.pretty_print() str += self.typeid.pretty_print()
if self.optional_typedecl is not None: if self.optional_typedecl is not None:
str += "," + self.optional_typedecl.pretty_print() str += ", "
str += self.optional_typedecl.pretty_print()
return str return str
@dataclass @dataclass
class LetfunDecl(ast_utils.Ast): class LetfunDecl(ast_utils.Ast):
ident: Ident ident: Ident
typedecl: Optional[List[Ident]] typedecl: Optional[Typedecl]
pterm: Pterm pterm: Pterm
def pretty_print(self):
str = f"letfun {self.ident.value}"
if self.typedecl is not None:
str += "("
str += self.typedecl.pretty_print()
# str += ", ".join(
# [
# t.pretty_print() if hasattr(t, "pretty_print") else str(t)
# for t in self.typedecl
# ]
# )
str += ")"
str += " = "
str += self.pterm.pretty_print()
str += "."
return str
@dataclass @dataclass
class LemmaGterm(ast_utils.Ast): class LemmaGterm(ast_utils.Ast):
@@ -229,8 +343,7 @@ class LemmaGterm(ast_utils.Ast):
lemma: Optional[Lemma] = None lemma: Optional[Lemma] = None
def pretty_print(self): def pretty_print(self):
str = "" 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()
return str return str
@@ -239,29 +352,53 @@ class LemmaGterm(ast_utils.Ast):
@dataclass @dataclass
class LemmaPublicVars(ast_utils.Ast): class LemmaPublicVars(ast_utils.Ast):
gterm: Gterm gterm: Gterm
public_vars: Optional[List] = None public_vars: List
lemma: Optional[Lemma] = None lemma: Optional[Lemma] = None
def pretty_print(self):
str = self.gterm.pretty_print()
str += "for { public_vars"
str += ",".join(
[
v.pretty_print() if hasattr(v, "pretty_print") else str(v)
for v in self.public_vars
]
)
str += "}"
if self.lemma is not None:
str += ";" + self.lemma.pretty_print()
return str
@dataclass @dataclass
class LemmaPublicVarsSecret(ast_utils.Ast): class LemmaPublicVarsSecret(ast_utils.Ast):
gterm: Gterm gterm: Gterm
secret: Ident secret: Ident
public_vars: Optional[List] = None public_vars: List
lemma: Optional[Lemma] = None lemma: Optional[Lemma] = None
def pretty_print(self):
str = self.gterm.pretty_print()
str += "for { secret"
str += self.secret.value
str += ",".join(
[
v.pretty_print() if hasattr(v, "pretty_print") else str(v)
for v in self.public_vars
]
)
str += "[real_or_random] }"
if self.lemma is not None:
str += ";" + self.lemma.pretty_print()
return str
@dataclass @dataclass
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):
str = "" return self.lemma.pretty_print()
if isinstance(self.lemma, LemmaGterm):
str += self.lemma.pretty_print()
else:
pass
return str
# def __post_init__(self): # def __post_init__(self):
# print(f"[constructor] Lemma: lemma={self.lemma}") # print(f"[constructor] Lemma: lemma={self.lemma}")
@@ -281,10 +418,10 @@ class LemmaDeclCore(ast_utils.Ast):
# lemma_decl_core: "lemma" [ typedecl ";"] lemma "." # lemma_decl_core: "lemma" [ typedecl ";"] lemma "."
def pretty_print(self): def pretty_print(self):
str = "" str = ""
str += "lemma" str += "lemma "
if self.typedecl is not None: if self.typedecl is not None:
str += self.typedecl.pretty_print() str += self.typedecl.pretty_print()
str += ";" str += ";\n "
str += self.lemma.pretty_print() str += self.lemma.pretty_print()
str += "." str += "."
return str return str
@@ -327,22 +464,40 @@ 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):
return self.query.pretty_print()
@dataclass @dataclass
class QueryAnnotation(ast_utils.Ast): class QueryAnnotation(ast_utils.Ast):
annotation: str annotation: str
def pretty_print(self):
return self.annotation
@dataclass @dataclass
class ReachableAnnotation(ast_utils.Ast): class ReachableAnnotation(ast_utils.Ast):
annotation: str annotation: str
def pretty_print(self):
return self.annotation
@dataclass @dataclass
class QueryDeclCore(ast_utils.Ast): class QueryDeclCore(ast_utils.Ast):
typedecl: Optional[Typedecl] typedecl: Optional[Typedecl]
query: Query query: Query
def pretty_print(self):
str = "query "
if self.typedecl is not None:
str += self.typedecl.pretty_print()
str += ";\n "
str += self.query.pretty_print()
str += "."
return str
@dataclass @dataclass
class QueryDecl(ast_utils.Ast): class QueryDecl(ast_utils.Ast):
@@ -350,7 +505,14 @@ class QueryDecl(ast_utils.Ast):
query_decl_core: QueryDeclCore query_decl_core: QueryDeclCore
def pretty_print(self): def pretty_print(self):
return "QueryDecl not implemented" str = ""
if self.query_decl_annotation is not None:
if isinstance(self.query_decl_annotation, QueryAnnotation):
str += f"@query {self.query_decl_annotation.pretty_print()}\n"
elif isinstance(self.query_decl_annotation, ReachableAnnotation):
str += f"@reachable {self.query_decl_annotation.pretty_print()}\n"
str += self.query_decl_core.pretty_print()
return str
@dataclass @dataclass
@@ -359,6 +521,20 @@ 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):
str = self.gterm.pretty_print()
if self.public_vars is not None:
str += "public_vars "
str += ", ".join(
[
v.pretty_print() if hasattr(v, "pretty_print") else str(v)
for v in self.public_vars
]
)
if self.query is not None:
str += ";" + self.query.pretty_print()
return str
@dataclass @dataclass
class QuerySecret(ast_utils.Ast): class QuerySecret(ast_utils.Ast):
@@ -366,16 +542,43 @@ 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):
str = "secret"
str += self.ident.pretty_print()
if self.public_vars is not None:
str += "public_vars "
str += ", ".join(
[
v.pretty_print() if hasattr(v, "pretty_print") else str(v)
for v in self.public_vars
]
)
if self.query is not None:
str += ";" + self.query.pretty_print()
return str
@dataclass @dataclass
class QueryPutBegin(ast_utils.Ast): 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):
str = "putbegin event :"
str += ", ".join(
[
e.pretty_print() if hasattr(e, "pretty_print") else str(e)
for e in self.event_list
]
)
if self.query is not None:
str += ";" + self.query.pretty_print()
return str
@dataclass @dataclass
class Decl(ast_utils.Ast): class Decl(ast_utils.Ast):
decl: LemmaDecl | QueryDecl decl: LemmaDecl | QueryDecl | TypeDecl | LetfunDecl
def pretty_print(self): def pretty_print(self):
return self.decl.pretty_print() return self.decl.pretty_print()
@@ -383,7 +586,7 @@ class Decl(ast_utils.Ast):
parser = Lark(""" parser = Lark("""
start: decl* start: decl*
decl: lemma_decl | query_decl | type_decl decl: lemma_decl | query_decl | type_decl | letfun_decl
_non_empty_seq{x}: x ("," x)* _non_empty_seq{x}: x ("," x)*
_maybe_empty_seq{x}: [ _non_empty_seq{x} ] _maybe_empty_seq{x}: [ _non_empty_seq{x} ]
IDENT: /[a-zA-Z][a-zA-Z0-9À-ÿ'_]*/ IDENT: /[a-zA-Z][a-zA-Z0-9À-ÿ'_]*/
@@ -655,35 +858,29 @@ def parse(input: str):
if __name__ == "__main__": if __name__ == "__main__":
parse(""" parse("""
type c. type c.
""") letfun foo (bar: bitstring, baz: int) = blub.
@lemma "secrecy: Adv can not learn shared secret key"
""" lemma kp:key_prec, skp:kem_sk_prec;
@lemma "secrecy: Adv can not learn shared secret key" attacker(trusted_key(kp)).
lemma kp:key_prec, skp:kem_sk_prec;
attacker(trusted_key(kp)).
@query "non-interruptability: Adv cannot start a responder session with the same key twice" @query "non-interruptability: Adv cannot start a responder session with the same key twice"
query ic1:InitConf_t, ic2:InitConf_t, ck:key, t1:time, t2:time; query ic1:InitConf_t, ic2:InitConf_t, ck:key, t1:time, t2:time;
event(ResponderSession(ic1, ck))@t1 && event(ResponderSession(ic2, ck))@t2 event(ResponderSession(ic1, ck))@t1 && event(ResponderSession(ic2, ck))@t2
==> t1 = t2. ==> t1 = t2.
@reachable "non-secrecy: The attacker can learn the value of a shared key" @reachable "non-secrecy: The attacker can learn the value of a shared key"
query k:key; query k:key;
attacker(prepare_key(k)) && attacker(k). attacker(prepare_key(k)) && attacker(k).
@lemma "secrecy: The adversary can learn a trusted kem pk only by using the reveal oracle" @lemma "secrecy: The adversary can learn a trusted kem pk only by using the reveal oracle"
lemma skp:kem_sk_prec; lemma skp:kem_sk_prec;
attacker(kem_pub(trusted_kem_sk(skp))) attacker(kem_pub(trusted_kem_sk(skp)))
==> event(RevealPk(kem_pub(trusted_kem_sk(skp)))). ==> event(RevealPk(kem_pub(trusted_kem_sk(skp)))).
@lemma "secrecy: Attacker knowledge of a kem sk implies the key is not trusted" @lemma "secrecy: Attacker knowledge of a kem sk implies the key is not trusted"
lemma k:kem_sk, kp:kem_sk_prec; lemma k:kem_sk, kp:kem_sk_prec;
attacker(prepare_kem_sk(k)) && attacker(k) ==> k <> trusted_kem_sk(kp). attacker(prepare_kem_sk(k)) && attacker(k) ==> k <> trusted_kem_sk(kp).
@lemma "asymmetric secrecy: Secure SSKI is sufficient for ck secrecy after trusted InitHello transmission from responder perspective" @lemma "asymmetric secrecy: Secure SSKI is sufficient for ck secrecy after trusted InitHello transmission from responder perspective"
lemma ck_ini:key, ck:key, any_psk:key, any_sskr:kem_sk, PSsski:kem_sk_prec, any_epki:kem_pk, any_epti:key, PSspti:seed_prec; lemma ck_ini:key, ck:key, any_psk:key, any_sskr:kem_sk, PSsski:kem_sk_prec, any_epki:kem_pk, any_epti:key, PSspti:seed_prec;
let secure_spki = kem_pub(trusted_kem_sk(PSsski)) in let secure_spki = kem_pub(trusted_kem_sk(PSsski)) in
let secure_spti = rng_key(trusted_seed(PSspti)) in let secure_spti = rng_key(trusted_seed(PSspti)) in
attacker(ck) attacker(ck)
&& event(OskOinit_hello(ck_ini, ck, any_psk, any_sskr, secure_spki, any_epki, any_epti, secure_spti)). && event(OskOinit_hello(ck_ini, ck, any_psk, any_sskr, secure_spki, any_epki, any_epti, secure_spti)).
""" """)