From a62338b2f1a9bf96183b9648464e1ea94812dd53 Mon Sep 17 00:00:00 2001 From: Will Xiao Date: Tue, 28 Jul 2026 10:20:56 +0000 Subject: [PATCH] Added functionality for set notation --- lf_toolkit/parse/set/__init__.py | 1 + lf_toolkit/parse/set/ast.py | 9 ++++++++- lf_toolkit/parse/set/grammar/ascii.lark | 5 ++++- lf_toolkit/parse/set/parser.py | 4 ++++ lf_toolkit/parse/set/printer.py | 8 +++++++- lf_toolkit/parse/set/transformer.py | 6 +++++- 6 files changed, 29 insertions(+), 4 deletions(-) diff --git a/lf_toolkit/parse/set/__init__.py b/lf_toolkit/parse/set/__init__.py index 446604a..589efe2 100644 --- a/lf_toolkit/parse/set/__init__.py +++ b/lf_toolkit/parse/set/__init__.py @@ -15,3 +15,4 @@ from .printer import UnicodePrinter from .transformer import SymPyBooleanTransformer from .transformer import SymPyTransformer +from .ast import SetNotation diff --git a/lf_toolkit/parse/set/ast.py b/lf_toolkit/parse/set/ast.py index 5372268..14a0bab 100644 --- a/lf_toolkit/parse/set/ast.py +++ b/lf_toolkit/parse/set/ast.py @@ -1,6 +1,7 @@ from abc import ABC from abc import abstractmethod -from dataclasses import dataclass +from dataclasses import dataclass, field +from typing import Tuple @dataclass @@ -58,6 +59,9 @@ class SymmetricDifference(BinaryOp): class Term(Set): value: str +@dataclass +class SetNotation(Set): + elements: Tuple[str, ...] = field(default_factory=tuple) class Universe(Set): pass @@ -80,6 +84,9 @@ def transform(self, node: Set): transformed_children.append(self.transform(node.left)) transformed_children.append(self.transform(node.right)) + if isinstance(node, SetNotation): + transformed_children.append(node.elements) + # Dispatch to the specific transformation method based on node type method_name = type(node).__name__ transformer = getattr(self, method_name, self.__unhandled_node) diff --git a/lf_toolkit/parse/set/grammar/ascii.lark b/lf_toolkit/parse/set/grammar/ascii.lark index 70a2b6e..0670952 100644 --- a/lf_toolkit/parse/set/grammar/ascii.lark +++ b/lf_toolkit/parse/set/grammar/ascii.lark @@ -22,7 +22,10 @@ start: expression ?group: "(" expression ")" -> group -?term: ID | universe +?term: ID | universe | set_notation + +set_notation: "{" (ELEMENT ("," ELEMENT)*)? "}" ID: /[A-Z]/ universe: "Ω" | "Omega" +ELEMENT: /[A-Za-z0-9]+/ diff --git a/lf_toolkit/parse/set/parser.py b/lf_toolkit/parse/set/parser.py index a062d41..50d03f6 100644 --- a/lf_toolkit/parse/set/parser.py +++ b/lf_toolkit/parse/set/parser.py @@ -17,6 +17,7 @@ from .ast import Term from .ast import Union from .ast import Universe +from .ast import SetNotation class ParseError(Exception): @@ -105,3 +106,6 @@ def universe(self, _): def group(self, items): return Group(items[1] if self.latex else items[0]) + + def set_notation(self, items): + return SetNotation(tuple(str(i) for i in items)) diff --git a/lf_toolkit/parse/set/printer.py b/lf_toolkit/parse/set/printer.py index 7818400..74cf6b8 100644 --- a/lf_toolkit/parse/set/printer.py +++ b/lf_toolkit/parse/set/printer.py @@ -1,5 +1,5 @@ from .ast import Set -from .ast import SetTransformer +from .ast import SetTransformer, SetNotation class LatexPrinter(SetTransformer): @@ -30,6 +30,9 @@ def Term(self, value): def Universe(self): return "\\Omega" + def SetNotation(self, elements): + return "\\{" + ",".join(elements) + "\\}" + class ASCIIPrinter(SetTransformer): def print(self, node: Set): @@ -59,6 +62,9 @@ def Term(self, value): def Universe(self): return "Omega" + def SetNotation(self, elements): + return "{" + ",".join(elements) + "}" + class UnicodePrinter(SetTransformer): def print(self, node: Set): diff --git a/lf_toolkit/parse/set/transformer.py b/lf_toolkit/parse/set/transformer.py index ecfaa7a..ff1994c 100644 --- a/lf_toolkit/parse/set/transformer.py +++ b/lf_toolkit/parse/set/transformer.py @@ -9,8 +9,9 @@ from sympy import Union from sympy import UniversalSet from sympy import Xor +from sympy import Integer -from .ast import SetTransformer +from .ast import SetTransformer, SetNotation class SymPyTransformer(SetTransformer): @@ -43,6 +44,9 @@ def Term(self, expr): def Universe(self): return UniversalSet + def SetNotation(self, elements): + return FiniteSet(*[Integer(e) if str(e).isdigit() else Symbol(e) for e in elements]) + class SymPyBooleanTransformer(SetTransformer): def __init__(self):