Compare commits
5
Commits
a735113466
...
4d23e8840e
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4d23e8840e
|
||
|
|
c64d626d1c
|
||
|
|
ecab1b74a4
|
||
|
|
0bbdf04621
|
||
|
|
939e5af4ce
|
No files matched your search
+78
-68
@@ -3,58 +3,34 @@ import re
|
||||
|
||||
HEADER = '''"""
|
||||
This file was generated by a script. Any manual changes might be overwritten.
|
||||
Please modify gen/ast.py instead and run gen/gen.py
|
||||
Please modify {defs_path} instead and run {gen_path}
|
||||
"""'''
|
||||
|
||||
SECTION_TEMPLATE = """{banner}
|
||||
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class {base}(ABC):
|
||||
location: Optional[Location] = None
|
||||
|
||||
@abstractmethod
|
||||
def accept(self, visitor: Visitor[T]) -> T: ...
|
||||
|
||||
class Visitor(ABC, Generic[T]):
|
||||
{visitor_methods}
|
||||
|
||||
|
||||
{classes}"""
|
||||
|
||||
TEMPLATE = """{header}
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Generic, Optional, TypeVar
|
||||
|
||||
from midas.ast.location import Location
|
||||
from midas.lexer.token import Token
|
||||
{imports}
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
##############
|
||||
# Statements #
|
||||
##############
|
||||
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class Stmt(ABC):
|
||||
location: Optional[Location] = None
|
||||
|
||||
@abstractmethod
|
||||
def accept(self, visitor: Visitor[T]) -> T: ...
|
||||
|
||||
class Visitor(ABC, Generic[T]):
|
||||
{stmt_visitor_methods}
|
||||
|
||||
|
||||
{statements}
|
||||
|
||||
|
||||
###############
|
||||
# Expressions #
|
||||
###############
|
||||
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class Expr(ABC):
|
||||
location: Optional[Location] = None
|
||||
|
||||
@abstractmethod
|
||||
def accept(self, visitor: Visitor[T]) -> T: ...
|
||||
|
||||
class Visitor(ABC, Generic[T]):
|
||||
{expr_visitor_methods}
|
||||
|
||||
|
||||
{expressions}
|
||||
{sections}
|
||||
"""
|
||||
|
||||
VISITOR_METHOD_TEMPLATE = """
|
||||
@@ -71,6 +47,16 @@ class {cls}({base}):
|
||||
return visitor.visit_{func_name}(self)
|
||||
"""
|
||||
|
||||
SECTION_REGEX = re.compile(
|
||||
r"^###>\s*(?P<base>[^\n]*?)\s*\|\s*(?P<name>[^\n]*?)(\s*\|\s*(?P<param>[^\n]*?))?\s*?\n(?P<body>.*?)\n###<$",
|
||||
re.MULTILINE | re.DOTALL,
|
||||
)
|
||||
|
||||
IMPORTS_REGEX = re.compile(
|
||||
r"^###>\s*Imports\s*?\n(?P<body>.*?)\n###<$",
|
||||
re.MULTILINE | re.DOTALL,
|
||||
)
|
||||
|
||||
|
||||
def snake_case(text: str) -> str:
|
||||
return re.sub(r"[A-Z]", lambda c: "_" + c.group().lower(), text).lower().strip("_")
|
||||
@@ -95,41 +81,65 @@ def make_class(name: str, cls: str, base: str):
|
||||
return cls_def.strip("\n")
|
||||
|
||||
|
||||
def generate(src: str):
|
||||
classes: list[str] = src.split("\n\n")
|
||||
stmt_visitor_methods: list[str] = []
|
||||
expr_visitor_methods: list[str] = []
|
||||
statements: list[str] = []
|
||||
expressions: list[str] = []
|
||||
def make_banner(text: str) -> str:
|
||||
middle: str = f"# {text} #"
|
||||
rule: str = "#" * len(middle)
|
||||
return "\n".join((rule, middle, rule))
|
||||
|
||||
for cls in classes:
|
||||
|
||||
def make_section(full_name: str, base: str, param: str, body: str) -> str:
|
||||
visitor_methods: list[str] = []
|
||||
classes: list[str] = []
|
||||
definitions: list[str] = body.strip("\n").split("\n\n\n")
|
||||
for cls in definitions:
|
||||
cls = cls.strip("\n")
|
||||
name: str = re.match("class (.*?):", cls).group(1) # type: ignore
|
||||
print(f"Processing {name}")
|
||||
if name.endswith("Stmt"):
|
||||
stmt_visitor_methods.append(make_visitor_method(name, "stmt"))
|
||||
statements.append(make_class(name, cls, "Stmt"))
|
||||
elif name.endswith("Expr"):
|
||||
expr_visitor_methods.append(make_visitor_method(name, "expr"))
|
||||
expressions.append(make_class(name, cls, "Expr"))
|
||||
visitor_methods.append(make_visitor_method(name, param))
|
||||
classes.append(make_class(name, cls, base))
|
||||
|
||||
return TEMPLATE.format(
|
||||
header=HEADER,
|
||||
stmt_visitor_methods="\n\n".join(stmt_visitor_methods),
|
||||
expr_visitor_methods="\n\n".join(expr_visitor_methods),
|
||||
statements="\n\n\n".join(statements),
|
||||
expressions="\n\n\n".join(expressions),
|
||||
return SECTION_TEMPLATE.format(
|
||||
banner=make_banner(full_name),
|
||||
base=base,
|
||||
visitor_methods="\n\n".join(visitor_methods),
|
||||
classes="\n\n\n".join(classes),
|
||||
)
|
||||
|
||||
|
||||
def generate(definitions_path: Path, out_path: Path):
|
||||
root_dir: Path = Path(__file__).parent.parent
|
||||
rel_path: Path = definitions_path.relative_to(root_dir)
|
||||
src: str = definitions_path.read_text()
|
||||
sections: list[str] = []
|
||||
|
||||
imports: str = ""
|
||||
if m := IMPORTS_REGEX.search(src):
|
||||
imports = m.group("body").strip("\n")
|
||||
|
||||
for section_m in SECTION_REGEX.finditer(src):
|
||||
full_name: str = section_m.group("name")
|
||||
base: str = section_m.group("base")
|
||||
param: str = section_m.group("param") or base.lower()
|
||||
body: str = section_m.group("body")
|
||||
sections.append(make_section(full_name, base, param, body))
|
||||
|
||||
result: str = TEMPLATE.format(
|
||||
header=HEADER.format(
|
||||
defs_path=rel_path,
|
||||
gen_path=Path(__file__).relative_to(root_dir),
|
||||
),
|
||||
imports=imports,
|
||||
sections="\n\n\n".join(sections),
|
||||
)
|
||||
out_path.write_text(result)
|
||||
|
||||
|
||||
def main():
|
||||
root: Path = Path(__file__).parent.parent
|
||||
in_path: Path = root / "gen" / "ast.py"
|
||||
out_path: Path = root / "midas" / "ast" / "midas.py"
|
||||
|
||||
src: str = in_path.read_text()
|
||||
generated: str = generate(src)
|
||||
out_path.write_text(generated)
|
||||
defs_dir: Path = root / "gen"
|
||||
ast_dir: Path = root / "midas" / "ast"
|
||||
generate(defs_dir / "midas.py", ast_dir / "midas.py")
|
||||
generate(defs_dir / "python.py", ast_dir / "python.py")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
+79
-41
@@ -1,72 +1,110 @@
|
||||
# type: ignore
|
||||
# ruff: disable[F821, F401]
|
||||
|
||||
###> Imports
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Generic, Optional, TypeVar
|
||||
|
||||
from midas.ast.location import Location
|
||||
from midas.lexer.token import Token
|
||||
|
||||
###<
|
||||
|
||||
|
||||
###> Stmt | Statements
|
||||
class SimpleTypeStmt:
|
||||
name: Token
|
||||
template: Optional[TemplateExpr]
|
||||
base: TypeExpr
|
||||
constraint: Optional[Expr]
|
||||
|
||||
class SimpleTypeExpr:
|
||||
name: Token
|
||||
optional: bool
|
||||
|
||||
class LogicalExpr:
|
||||
left: Expr
|
||||
operator: Token
|
||||
right: Expr
|
||||
|
||||
class BinaryExpr:
|
||||
left: Expr
|
||||
operator: Token
|
||||
right: Expr
|
||||
|
||||
class UnaryExpr:
|
||||
operator: Token
|
||||
right: Expr
|
||||
|
||||
class GetExpr:
|
||||
expr: Expr
|
||||
name: Token
|
||||
|
||||
class VariableExpr:
|
||||
name: Token
|
||||
|
||||
class GroupingExpr:
|
||||
expr: Expr
|
||||
|
||||
class LiteralExpr:
|
||||
value: Any
|
||||
|
||||
class WildcardExpr:
|
||||
token: Token
|
||||
|
||||
class TemplateExpr:
|
||||
type: TypeExpr
|
||||
|
||||
class TypeExpr:
|
||||
name: Token
|
||||
template: Optional[TemplateExpr]
|
||||
optional: bool
|
||||
|
||||
class ComplexTypeStmt:
|
||||
name: Token
|
||||
template: Optional[TemplateExpr]
|
||||
properties: list[PropertyStmt]
|
||||
|
||||
|
||||
class PropertyStmt:
|
||||
name: Token
|
||||
type: TypeExpr
|
||||
constraint: Optional[Expr]
|
||||
|
||||
|
||||
class ExtendStmt:
|
||||
type: TypeExpr
|
||||
operations: list[OpStmt]
|
||||
|
||||
|
||||
class OpStmt:
|
||||
name: Token
|
||||
operand: TypeExpr
|
||||
result: TypeExpr
|
||||
|
||||
|
||||
class PredicateStmt:
|
||||
name: Token
|
||||
subject: Token
|
||||
type: TypeExpr
|
||||
condition: Expr
|
||||
|
||||
|
||||
###<
|
||||
|
||||
|
||||
###> Expr | Expressions
|
||||
class SimpleTypeExpr:
|
||||
name: Token
|
||||
optional: bool
|
||||
|
||||
|
||||
class LogicalExpr:
|
||||
left: Expr
|
||||
operator: Token
|
||||
right: Expr
|
||||
|
||||
|
||||
class BinaryExpr:
|
||||
left: Expr
|
||||
operator: Token
|
||||
right: Expr
|
||||
|
||||
|
||||
class UnaryExpr:
|
||||
operator: Token
|
||||
right: Expr
|
||||
|
||||
|
||||
class GetExpr:
|
||||
expr: Expr
|
||||
name: Token
|
||||
|
||||
|
||||
class VariableExpr:
|
||||
name: Token
|
||||
|
||||
|
||||
class GroupingExpr:
|
||||
expr: Expr
|
||||
|
||||
|
||||
class LiteralExpr:
|
||||
value: Any
|
||||
|
||||
|
||||
class WildcardExpr:
|
||||
token: Token
|
||||
|
||||
|
||||
class TemplateExpr:
|
||||
type: TypeExpr
|
||||
|
||||
|
||||
class TypeExpr:
|
||||
name: Token
|
||||
template: Optional[TemplateExpr]
|
||||
optional: bool
|
||||
|
||||
|
||||
###<
|
||||
+112
@@ -0,0 +1,112 @@
|
||||
# type: ignore
|
||||
# ruff: disable[F821, F401]
|
||||
|
||||
###> Imports
|
||||
import ast
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Generic, Optional, TypeVar
|
||||
|
||||
from midas.ast.location import Location
|
||||
|
||||
###<
|
||||
|
||||
|
||||
###> MidasType | Type annotations | node
|
||||
class BaseType:
|
||||
base: str
|
||||
param: Optional[MidasType]
|
||||
|
||||
|
||||
class ConstraintType:
|
||||
type: MidasType
|
||||
constraint: ast.expr
|
||||
|
||||
|
||||
class FrameColumn:
|
||||
name: Optional[str]
|
||||
type: Optional[MidasType]
|
||||
|
||||
|
||||
class FrameType:
|
||||
columns: list[FrameColumn]
|
||||
|
||||
|
||||
###<
|
||||
|
||||
|
||||
###> Stmt | Statements
|
||||
class ExpressionStmt:
|
||||
expr: Expr
|
||||
|
||||
|
||||
class Function:
|
||||
name: str
|
||||
posonlyargs: list[Argument]
|
||||
args: list[Argument]
|
||||
kwonlyargs: list[Argument]
|
||||
returns: Optional[MidasType]
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class Argument:
|
||||
location: Optional[Location] = None
|
||||
name: Optional[str]
|
||||
type: Optional[MidasType]
|
||||
|
||||
|
||||
class TypeAssign:
|
||||
name: str
|
||||
type: MidasType
|
||||
|
||||
|
||||
###<
|
||||
|
||||
|
||||
###> Expr | Expressions
|
||||
class AssignExpr:
|
||||
name: str
|
||||
value: Expr
|
||||
|
||||
|
||||
class BinaryExpr:
|
||||
left: Expr
|
||||
operator: ast.operator
|
||||
right: Expr
|
||||
|
||||
|
||||
class UnaryExpr:
|
||||
operator: ast.unaryop
|
||||
right: Expr
|
||||
|
||||
|
||||
class CallExpr:
|
||||
callee: Expr
|
||||
arguments: list[Expr]
|
||||
|
||||
|
||||
class GetExpr:
|
||||
object: Expr
|
||||
name: str
|
||||
|
||||
|
||||
class LiteralExpr:
|
||||
value: Any
|
||||
|
||||
|
||||
class VariableExpr:
|
||||
name: str
|
||||
|
||||
|
||||
class LogicalExpr:
|
||||
left: Expr
|
||||
operator: ast.boolop
|
||||
right: Expr
|
||||
|
||||
|
||||
class SetExpr:
|
||||
object: Expr
|
||||
name: str
|
||||
value: Expr
|
||||
|
||||
|
||||
###<
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
"""
|
||||
This file was generated by a script. Any manual changes might be overwritten.
|
||||
Please modify gen/ast.py instead and run gen/gen.py
|
||||
Please modify gen/midas.py instead and run gen/gen.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
+120
-17
@@ -350,7 +350,12 @@ class MidasPrinter(m.Expr.Visitor[str], m.Stmt.Visitor[str]):
|
||||
return f"{expr.name.lexeme}{template}{'?' if expr.optional else ''}"
|
||||
|
||||
|
||||
class PythonAstPrinter(AstPrinter, p.Expr.Visitor[None]):
|
||||
class PythonAstPrinter(
|
||||
AstPrinter,
|
||||
p.MidasType.Visitor[None],
|
||||
p.Stmt.Visitor[None],
|
||||
p.Expr.Visitor[None],
|
||||
):
|
||||
def visit_base_type(self, node: p.BaseType) -> None:
|
||||
self._write_line("BaseType")
|
||||
with self._child_level():
|
||||
@@ -382,39 +387,137 @@ class PythonAstPrinter(AstPrinter, p.Expr.Visitor[None]):
|
||||
self._mark_last()
|
||||
col.accept(self)
|
||||
|
||||
def visit_function(self, node: p.Function) -> None:
|
||||
def visit_expression_stmt(self, stmt: p.ExpressionStmt) -> None:
|
||||
stmt.expr.accept(self)
|
||||
|
||||
def visit_function(self, stmt: p.Function) -> None:
|
||||
self._write_line("Function")
|
||||
with self._child_level():
|
||||
self._write_line(f"name: {node.name}")
|
||||
self._write_line(f"name: {stmt.name}")
|
||||
|
||||
self._write_line("posonlyargs")
|
||||
with self._child_level():
|
||||
for i, arg in enumerate(node.posonlyargs):
|
||||
for i, arg in enumerate(stmt.posonlyargs):
|
||||
self._idx = i
|
||||
if i == len(node.posonlyargs) - 1:
|
||||
if i == len(stmt.posonlyargs) - 1:
|
||||
self._mark_last()
|
||||
arg.accept(self)
|
||||
self._print_argument(arg)
|
||||
|
||||
self._write_line("args")
|
||||
with self._child_level():
|
||||
for i, arg in enumerate(node.args):
|
||||
for i, arg in enumerate(stmt.args):
|
||||
self._idx = i
|
||||
if i == len(node.args) - 1:
|
||||
if i == len(stmt.args) - 1:
|
||||
self._mark_last()
|
||||
arg.accept(self)
|
||||
self._print_argument(arg)
|
||||
|
||||
self._write_line("kwonlyargs")
|
||||
with self._child_level():
|
||||
for i, arg in enumerate(node.kwonlyargs):
|
||||
for i, arg in enumerate(stmt.kwonlyargs):
|
||||
self._idx = i
|
||||
if i == len(node.kwonlyargs) - 1:
|
||||
if i == len(stmt.kwonlyargs) - 1:
|
||||
self._mark_last()
|
||||
self._print_argument(arg)
|
||||
|
||||
self._write_optional_child("returns", stmt.returns, last=True)
|
||||
|
||||
def _print_argument(self, arg: p.Function.Argument) -> None:
|
||||
self._write_line("FunctionArgument")
|
||||
with self._child_level():
|
||||
self._write_line(f"name: {arg.name}")
|
||||
self._write_optional_child("type", arg.type, last=True)
|
||||
|
||||
def visit_type_assign(self, stmt: p.TypeAssign) -> None:
|
||||
self._write_line("TypeAssign")
|
||||
with self._child_level():
|
||||
self._write_line(f"name: {stmt.name}")
|
||||
self._write_line("type", last=True)
|
||||
with self._child_level(single=True):
|
||||
stmt.type.accept(self)
|
||||
|
||||
def visit_assign_expr(self, expr: p.AssignExpr) -> None:
|
||||
self._write_line("AssignExpr")
|
||||
with self._child_level():
|
||||
self._write_line(f"name: {expr.name}")
|
||||
self._write_line("value", last=True)
|
||||
with self._child_level(single=True):
|
||||
expr.value.accept(self)
|
||||
|
||||
def visit_binary_expr(self, expr: p.BinaryExpr) -> None:
|
||||
self._write_line("BinaryExpr")
|
||||
with self._child_level():
|
||||
self._write_line("left")
|
||||
with self._child_level(single=True):
|
||||
expr.left.accept(self)
|
||||
|
||||
self._write_line(f"operator: {expr.operator.__class__.__name__}")
|
||||
|
||||
self._write_line("right", last=True)
|
||||
with self._child_level(single=True):
|
||||
expr.right.accept(self)
|
||||
|
||||
def visit_unary_expr(self, expr: p.UnaryExpr) -> None:
|
||||
self._write_line("UnaryExpr")
|
||||
with self._child_level():
|
||||
self._write_line(f"operator: {expr.operator.__class__.__name__}")
|
||||
|
||||
self._write_line("right", last=True)
|
||||
with self._child_level(single=True):
|
||||
expr.right.accept(self)
|
||||
|
||||
def visit_call_expr(self, expr: p.CallExpr) -> None:
|
||||
self._write_line("CallExpr")
|
||||
with self._child_level():
|
||||
self._write_line("callee")
|
||||
with self._child_level(single=True):
|
||||
expr.callee.accept(self)
|
||||
|
||||
self._write_line("arguments", last=True)
|
||||
with self._child_level():
|
||||
for i, arg in enumerate(expr.arguments):
|
||||
self._idx = i
|
||||
if i == len(expr.arguments) - 1:
|
||||
self._mark_last()
|
||||
arg.accept(self)
|
||||
|
||||
self._write_optional_child("returns", node.returns, last=True)
|
||||
|
||||
def visit_function_argument(self, node: p.FunctionArgument) -> None:
|
||||
self._write_line("FunctionArgument")
|
||||
def visit_get_expr(self, expr: p.GetExpr) -> None:
|
||||
self._write_line("GetExpr")
|
||||
with self._child_level():
|
||||
self._write_line(f"name: {node.name}")
|
||||
self._write_optional_child("type", node.type, last=True)
|
||||
self._write_line("object")
|
||||
with self._child_level(single=True):
|
||||
expr.object.accept(self)
|
||||
self._write_line(f"name: {expr.name}", last=True)
|
||||
|
||||
def visit_literal_expr(self, expr: p.LiteralExpr) -> None:
|
||||
self._write_line("LiteralExpr")
|
||||
with self._child_level(single=True):
|
||||
self._write_line(f"value: {expr.value}")
|
||||
|
||||
def visit_variable_expr(self, expr: p.VariableExpr) -> None:
|
||||
self._write_line("VariableExpr")
|
||||
with self._child_level(single=True):
|
||||
self._write_line(f"name: {expr.name}")
|
||||
|
||||
def visit_logical_expr(self, expr: p.LogicalExpr) -> None:
|
||||
self._write_line("LogicalExpr")
|
||||
with self._child_level():
|
||||
self._write_line("left")
|
||||
with self._child_level(single=True):
|
||||
expr.left.accept(self)
|
||||
|
||||
self._write_line(f"operator: {expr.operator.__class__.__name__}")
|
||||
|
||||
self._write_line("right", last=True)
|
||||
with self._child_level(single=True):
|
||||
expr.right.accept(self)
|
||||
|
||||
def visit_set_expr(self, expr: p.SetExpr) -> None:
|
||||
self._write_line("SetExpr")
|
||||
with self._child_level():
|
||||
self._write_line("object")
|
||||
with self._child_level(single=True):
|
||||
expr.object.accept(self)
|
||||
self._write_line(f"name: {expr.name}")
|
||||
self._write_line("value", last=True)
|
||||
with self._child_level(single=True):
|
||||
expr.value.accept(self)
|
||||
+185
-27
@@ -1,17 +1,26 @@
|
||||
"""
|
||||
This file was generated by a script. Any manual changes might be overwritten.
|
||||
Please modify gen/python.py instead and run gen/gen.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
import ast
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import Generic, Optional, TypeVar
|
||||
from typing import Any, Generic, Optional, TypeVar
|
||||
|
||||
from midas.ast.location import Location
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
####################
|
||||
# Type annotations #
|
||||
####################
|
||||
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class Expr(ABC):
|
||||
class MidasType(ABC):
|
||||
location: Optional[Location] = None
|
||||
|
||||
@abstractmethod
|
||||
@@ -30,24 +39,13 @@ class Expr(ABC):
|
||||
@abstractmethod
|
||||
def visit_frame_type(self, node: FrameType) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_function(self, node: Function) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_function_argument(self, node: FunctionArgument) -> T: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MidasType(Expr):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BaseType(MidasType):
|
||||
base: str
|
||||
param: Optional[MidasType]
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
def accept(self, visitor: MidasType.Visitor[T]) -> T:
|
||||
return visitor.visit_base_type(self)
|
||||
|
||||
|
||||
@@ -56,7 +54,7 @@ class ConstraintType(MidasType):
|
||||
type: MidasType
|
||||
constraint: ast.expr
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
def accept(self, visitor: MidasType.Visitor[T]) -> T:
|
||||
return visitor.visit_constraint_type(self)
|
||||
|
||||
|
||||
@@ -65,7 +63,7 @@ class FrameColumn(MidasType):
|
||||
name: Optional[str]
|
||||
type: Optional[MidasType]
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
def accept(self, visitor: MidasType.Visitor[T]) -> T:
|
||||
return visitor.visit_frame_column(self)
|
||||
|
||||
|
||||
@@ -73,26 +71,186 @@ class FrameColumn(MidasType):
|
||||
class FrameType(MidasType):
|
||||
columns: list[FrameColumn]
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
def accept(self, visitor: MidasType.Visitor[T]) -> T:
|
||||
return visitor.visit_frame_type(self)
|
||||
|
||||
|
||||
##############
|
||||
# Statements #
|
||||
##############
|
||||
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class Stmt(ABC):
|
||||
location: Optional[Location] = None
|
||||
|
||||
@abstractmethod
|
||||
def accept(self, visitor: Visitor[T]) -> T: ...
|
||||
|
||||
class Visitor(ABC, Generic[T]):
|
||||
@abstractmethod
|
||||
def visit_expression_stmt(self, stmt: ExpressionStmt) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_function(self, stmt: Function) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_type_assign(self, stmt: TypeAssign) -> T: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Function(Expr):
|
||||
class ExpressionStmt(Stmt):
|
||||
expr: Expr
|
||||
|
||||
def accept(self, visitor: Stmt.Visitor[T]) -> T:
|
||||
return visitor.visit_expression_stmt(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Function(Stmt):
|
||||
name: str
|
||||
posonlyargs: list[FunctionArgument]
|
||||
args: list[FunctionArgument]
|
||||
kwonlyargs: list[FunctionArgument]
|
||||
posonlyargs: list[Argument]
|
||||
args: list[Argument]
|
||||
kwonlyargs: list[Argument]
|
||||
returns: Optional[MidasType]
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class Argument:
|
||||
location: Optional[Location] = None
|
||||
name: Optional[str]
|
||||
type: Optional[MidasType]
|
||||
|
||||
def accept(self, visitor: Stmt.Visitor[T]) -> T:
|
||||
return visitor.visit_function(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FunctionArgument(Expr):
|
||||
name: Optional[str]
|
||||
type: Optional[MidasType]
|
||||
class TypeAssign(Stmt):
|
||||
name: str
|
||||
type: MidasType
|
||||
|
||||
def accept(self, visitor: Stmt.Visitor[T]) -> T:
|
||||
return visitor.visit_type_assign(self)
|
||||
|
||||
|
||||
###############
|
||||
# Expressions #
|
||||
###############
|
||||
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class Expr(ABC):
|
||||
location: Optional[Location] = None
|
||||
|
||||
@abstractmethod
|
||||
def accept(self, visitor: Visitor[T]) -> T: ...
|
||||
|
||||
class Visitor(ABC, Generic[T]):
|
||||
@abstractmethod
|
||||
def visit_assign_expr(self, expr: AssignExpr) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_binary_expr(self, expr: BinaryExpr) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_unary_expr(self, expr: UnaryExpr) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_call_expr(self, expr: CallExpr) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_get_expr(self, expr: GetExpr) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_literal_expr(self, expr: LiteralExpr) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_variable_expr(self, expr: VariableExpr) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_logical_expr(self, expr: LogicalExpr) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_set_expr(self, expr: SetExpr) -> T: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AssignExpr(Expr):
|
||||
name: str
|
||||
value: Expr
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_function_argument(self)
|
||||
return visitor.visit_assign_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BinaryExpr(Expr):
|
||||
left: Expr
|
||||
operator: ast.operator
|
||||
right: Expr
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_binary_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class UnaryExpr(Expr):
|
||||
operator: ast.unaryop
|
||||
right: Expr
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_unary_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CallExpr(Expr):
|
||||
callee: Expr
|
||||
arguments: list[Expr]
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_call_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GetExpr(Expr):
|
||||
object: Expr
|
||||
name: str
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_get_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LiteralExpr(Expr):
|
||||
value: Any
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_literal_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class VariableExpr(Expr):
|
||||
name: str
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_variable_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LogicalExpr(Expr):
|
||||
left: Expr
|
||||
operator: ast.boolop
|
||||
right: Expr
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_logical_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SetExpr(Expr):
|
||||
object: Expr
|
||||
name: str
|
||||
value: Expr
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_set_expr(self)
|
||||
+69
-26
@@ -1,15 +1,20 @@
|
||||
import ast
|
||||
from typing import Any, Optional
|
||||
from typing import Optional
|
||||
|
||||
from midas.ast.location import Location
|
||||
|
||||
from midas.ast.python import (
|
||||
AssignExpr,
|
||||
BaseType,
|
||||
ConstraintType,
|
||||
Expr,
|
||||
ExpressionStmt,
|
||||
FrameColumn,
|
||||
FrameType,
|
||||
Function,
|
||||
FunctionArgument,
|
||||
MidasType,
|
||||
Stmt,
|
||||
TypeAssign,
|
||||
)
|
||||
|
||||
|
||||
@@ -24,33 +29,66 @@ class UnsupportedSyntaxError(Exception):
|
||||
)
|
||||
|
||||
|
||||
class PythonParser(ast.NodeVisitor):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
class PythonParser:
|
||||
def parse_module(self, node: ast.Module) -> list[Stmt]:
|
||||
statements: list[Stmt] = []
|
||||
for stmt in node.body:
|
||||
parsed: None | Stmt | list[Stmt] = self.parse_stmt(stmt)
|
||||
if isinstance(parsed, Stmt):
|
||||
statements.append(parsed)
|
||||
elif parsed is not None:
|
||||
statements.extend(parsed)
|
||||
return statements
|
||||
|
||||
self.annotations: list[tuple[str, Optional[MidasType]]] = []
|
||||
self.functions: list[Function] = []
|
||||
|
||||
def visit_AnnAssign(self, node: ast.AnnAssign) -> Any:
|
||||
def parse_stmt(self, node: ast.stmt) -> None | Stmt | list[Stmt]:
|
||||
match node:
|
||||
case ast.AnnAssign(
|
||||
target=ast.Name(id=target), annotation=annotation, simple=1
|
||||
):
|
||||
self.annotations.append(
|
||||
(target, self._parse_type(annotation, root=True))
|
||||
)
|
||||
case ast.AnnAssign():
|
||||
return self.parse_annotation_assign(node)
|
||||
|
||||
case ast.FunctionDef():
|
||||
return self.parse_function(node)
|
||||
|
||||
case _:
|
||||
print(f"Unsupported assignment: {ast.unparse(node)}")
|
||||
return None
|
||||
|
||||
def parse_annotation_assign(self, node: ast.AnnAssign) -> list[Stmt]:
|
||||
statements: list[Stmt] = []
|
||||
loc: Location = Location.from_ast(node)
|
||||
match node:
|
||||
case ast.AnnAssign(
|
||||
target=ast.Name(id=target),
|
||||
annotation=annotation,
|
||||
value=value,
|
||||
simple=1,
|
||||
):
|
||||
type = self._parse_type(annotation, root=True)
|
||||
if type is not None:
|
||||
statements.append(
|
||||
TypeAssign(
|
||||
location=loc,
|
||||
name=target,
|
||||
type=type,
|
||||
)
|
||||
)
|
||||
|
||||
if value is not None:
|
||||
parsed_value: Expr = self.parse_expr(value)
|
||||
statements.append(
|
||||
ExpressionStmt(
|
||||
location=loc,
|
||||
expr=AssignExpr(
|
||||
location=loc,
|
||||
name=target,
|
||||
value=parsed_value,
|
||||
),
|
||||
)
|
||||
)
|
||||
case _:
|
||||
print(f"Unsupported annotation: {ast.unparse(node)}")
|
||||
return statements
|
||||
|
||||
def visit_FunctionDef(self, node: ast.FunctionDef) -> Any:
|
||||
self.functions.append(self._parse_function(node))
|
||||
|
||||
# Call visit on children to process body
|
||||
# TODO: scope the resulting nodes to the function
|
||||
self.generic_visit(node)
|
||||
|
||||
def _parse_function(self, node: ast.FunctionDef) -> Function:
|
||||
def parse_function(self, node: ast.FunctionDef) -> Function:
|
||||
loc: Location = Location.from_ast(node)
|
||||
match node:
|
||||
case ast.FunctionDef(
|
||||
@@ -63,7 +101,7 @@ class PythonParser(ast.NodeVisitor):
|
||||
returns=returns,
|
||||
):
|
||||
|
||||
def parse_args(args_list: list[ast.arg]) -> list[FunctionArgument]:
|
||||
def parse_args(args_list: list[ast.arg]) -> list[Function.Argument]:
|
||||
return [self._parse_function_argument(arg) for arg in args_list]
|
||||
|
||||
return Function(
|
||||
@@ -74,14 +112,16 @@ class PythonParser(ast.NodeVisitor):
|
||||
kwonlyargs=parse_args(kwonlyargs),
|
||||
returns=self._parse_type(returns) if returns is not None else None,
|
||||
)
|
||||
case _:
|
||||
print(f"Unsupported function definition: {ast.unparse(node)}")
|
||||
|
||||
def _parse_function_argument(self, arg: ast.arg) -> FunctionArgument:
|
||||
def _parse_function_argument(self, arg: ast.arg) -> Function.Argument:
|
||||
loc: Location = Location.from_ast(arg)
|
||||
name: str = arg.arg
|
||||
type: Optional[MidasType] = None
|
||||
if arg.annotation is not None:
|
||||
type = self._parse_type(arg.annotation)
|
||||
return FunctionArgument(
|
||||
return Function.Argument(
|
||||
location=loc,
|
||||
name=name,
|
||||
type=type,
|
||||
@@ -186,3 +226,6 @@ class PythonParser(ast.NodeVisitor):
|
||||
|
||||
case _:
|
||||
raise UnsupportedSyntaxError(column)
|
||||
|
||||
def parse_expr(self, node: ast.expr) -> Expr:
|
||||
raise NotImplementedError()
|
||||
Reference in new issue
Block a user