Compare commits
2
Commits
bbd0e3ae8d
...
0b3f33d7fe
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0b3f33d7fe
|
||
|
|
8a9b4f3989
|
No files matched your search
+11
-5
@@ -59,21 +59,27 @@ class TypeAssign:
|
||||
type: MidasType
|
||||
|
||||
|
||||
class AssignStmt:
|
||||
targets: list[Expr]
|
||||
value: Expr
|
||||
|
||||
|
||||
###<
|
||||
|
||||
|
||||
###> Expr | Expressions
|
||||
class AssignExpr:
|
||||
name: str
|
||||
value: Expr
|
||||
|
||||
|
||||
class BinaryExpr:
|
||||
left: Expr
|
||||
operator: ast.operator
|
||||
right: Expr
|
||||
|
||||
|
||||
class CompareExpr:
|
||||
left: Expr
|
||||
operator: ast.cmpop
|
||||
right: Expr
|
||||
|
||||
|
||||
class UnaryExpr:
|
||||
operator: ast.unaryop
|
||||
right: Expr
|
||||
|
||||
+34
-5
@@ -435,13 +435,19 @@ class PythonAstPrinter(
|
||||
with self._child_level(single=True):
|
||||
stmt.type.accept(self)
|
||||
|
||||
def visit_assign_expr(self, expr: p.AssignExpr) -> None:
|
||||
self._write_line("AssignExpr")
|
||||
def visit_assign_stmt(self, stmt: p.AssignStmt) -> None:
|
||||
self._write_line("AssignStmt")
|
||||
with self._child_level():
|
||||
self._write_line(f"name: {expr.name}")
|
||||
self._write_line("targets")
|
||||
with self._child_level():
|
||||
for i, target in enumerate(stmt.targets):
|
||||
self._idx = i
|
||||
if i == len(stmt.targets) - 1:
|
||||
self._mark_last()
|
||||
target.accept(self)
|
||||
self._write_line("value", last=True)
|
||||
with self._child_level(single=True):
|
||||
expr.value.accept(self)
|
||||
stmt.value.accept(self)
|
||||
|
||||
def visit_binary_expr(self, expr: p.BinaryExpr) -> None:
|
||||
self._write_line("BinaryExpr")
|
||||
@@ -456,6 +462,19 @@ class PythonAstPrinter(
|
||||
with self._child_level(single=True):
|
||||
expr.right.accept(self)
|
||||
|
||||
def visit_compare_expr(self, expr: p.CompareExpr) -> None:
|
||||
self._write_line("CompareExpr")
|
||||
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():
|
||||
@@ -472,7 +491,7 @@ class PythonAstPrinter(
|
||||
with self._child_level(single=True):
|
||||
expr.callee.accept(self)
|
||||
|
||||
self._write_line("arguments", last=True)
|
||||
self._write_line("arguments")
|
||||
with self._child_level():
|
||||
for i, arg in enumerate(expr.arguments):
|
||||
self._idx = i
|
||||
@@ -480,6 +499,16 @@ class PythonAstPrinter(
|
||||
self._mark_last()
|
||||
arg.accept(self)
|
||||
|
||||
self._write_line("keywords", last=True)
|
||||
with self._child_level():
|
||||
for i, (name, arg) in enumerate(expr.keywords.items()):
|
||||
self._idx = i
|
||||
if i == len(expr.keywords) - 1:
|
||||
self._mark_last()
|
||||
self._write_line(name)
|
||||
with self._child_level(single=True):
|
||||
arg.accept(self)
|
||||
|
||||
def visit_get_expr(self, expr: p.GetExpr) -> None:
|
||||
self._write_line("GetExpr")
|
||||
with self._child_level():
|
||||
|
||||
+25
-11
@@ -97,6 +97,9 @@ class Stmt(ABC):
|
||||
@abstractmethod
|
||||
def visit_type_assign(self, stmt: TypeAssign) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_assign_stmt(self, stmt: AssignStmt) -> T: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ExpressionStmt(Stmt):
|
||||
@@ -133,6 +136,15 @@ class TypeAssign(Stmt):
|
||||
return visitor.visit_type_assign(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AssignStmt(Stmt):
|
||||
targets: list[Expr]
|
||||
value: Expr
|
||||
|
||||
def accept(self, visitor: Stmt.Visitor[T]) -> T:
|
||||
return visitor.visit_assign_stmt(self)
|
||||
|
||||
|
||||
###############
|
||||
# Expressions #
|
||||
###############
|
||||
@@ -147,10 +159,10 @@ class Expr(ABC):
|
||||
|
||||
class Visitor(ABC, Generic[T]):
|
||||
@abstractmethod
|
||||
def visit_assign_expr(self, expr: AssignExpr) -> T: ...
|
||||
def visit_binary_expr(self, expr: BinaryExpr) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_binary_expr(self, expr: BinaryExpr) -> T: ...
|
||||
def visit_compare_expr(self, expr: CompareExpr) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_unary_expr(self, expr: UnaryExpr) -> T: ...
|
||||
@@ -174,15 +186,6 @@ class Expr(ABC):
|
||||
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_assign_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BinaryExpr(Expr):
|
||||
left: Expr
|
||||
@@ -193,6 +196,16 @@ class BinaryExpr(Expr):
|
||||
return visitor.visit_binary_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CompareExpr(Expr):
|
||||
left: Expr
|
||||
operator: ast.cmpop
|
||||
right: Expr
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_compare_expr(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class UnaryExpr(Expr):
|
||||
operator: ast.unaryop
|
||||
@@ -206,6 +219,7 @@ class UnaryExpr(Expr):
|
||||
class CallExpr(Expr):
|
||||
callee: Expr
|
||||
arguments: list[Expr]
|
||||
keywords: dict[str, Expr]
|
||||
|
||||
def accept(self, visitor: Expr.Visitor[T]) -> T:
|
||||
return visitor.visit_call_expr(self)
|
||||
|
||||
@@ -151,10 +151,12 @@ class PythonHighlighter(
|
||||
def visit_type_assign(self, stmt: p.TypeAssign) -> None:
|
||||
stmt.type.accept(self)
|
||||
|
||||
def visit_assign_expr(self, expr: p.AssignExpr) -> None: ...
|
||||
def visit_assign_stmt(self, stmt: p.AssignStmt) -> None: ...
|
||||
|
||||
def visit_binary_expr(self, expr: p.BinaryExpr) -> None: ...
|
||||
|
||||
def visit_compare_expr(self, expr: p.CompareExpr) -> None: ...
|
||||
|
||||
def visit_unary_expr(self, expr: p.UnaryExpr) -> None: ...
|
||||
|
||||
def visit_call_expr(self, expr: p.CallExpr) -> None: ...
|
||||
|
||||
+128
-16
@@ -4,17 +4,25 @@ from typing import Optional
|
||||
from midas.ast.location import Location
|
||||
|
||||
from midas.ast.python import (
|
||||
AssignExpr,
|
||||
AssignStmt,
|
||||
BaseType,
|
||||
BinaryExpr,
|
||||
CallExpr,
|
||||
CompareExpr,
|
||||
ConstraintType,
|
||||
Expr,
|
||||
ExpressionStmt,
|
||||
FrameColumn,
|
||||
FrameType,
|
||||
Function,
|
||||
GetExpr,
|
||||
LiteralExpr,
|
||||
LogicalExpr,
|
||||
MidasType,
|
||||
Stmt,
|
||||
TypeAssign,
|
||||
UnaryExpr,
|
||||
VariableExpr,
|
||||
)
|
||||
|
||||
|
||||
@@ -33,11 +41,15 @@ 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)
|
||||
try:
|
||||
parsed: None | Stmt | list[Stmt] = self.parse_stmt(stmt)
|
||||
if isinstance(parsed, Stmt):
|
||||
statements.append(parsed)
|
||||
elif parsed is not None:
|
||||
statements.extend(parsed)
|
||||
except UnsupportedSyntaxError as e:
|
||||
print(f"{e}, skipping")
|
||||
continue
|
||||
return statements
|
||||
|
||||
def parse_stmt(self, node: ast.stmt) -> None | Stmt | list[Stmt]:
|
||||
@@ -45,11 +57,17 @@ class PythonParser:
|
||||
case ast.AnnAssign():
|
||||
return self.parse_annotation_assign(node)
|
||||
|
||||
case ast.Assign():
|
||||
return self.parse_assign(node)
|
||||
|
||||
case ast.FunctionDef():
|
||||
return self.parse_function(node)
|
||||
|
||||
case ast.Expr(value=expr):
|
||||
return ExpressionStmt(expr=self.parse_expr(expr))
|
||||
|
||||
case _:
|
||||
print(f"Unsupported assignment: {ast.unparse(node)}")
|
||||
print(f"Unsupported statement: {ast.unparse(node)}")
|
||||
return None
|
||||
|
||||
def parse_annotation_assign(self, node: ast.AnnAssign) -> list[Stmt]:
|
||||
@@ -73,21 +91,32 @@ class PythonParser:
|
||||
)
|
||||
|
||||
if value is not None:
|
||||
parsed_value: Expr = self.parse_expr(value)
|
||||
statements.append(
|
||||
ExpressionStmt(
|
||||
AssignStmt(
|
||||
location=loc,
|
||||
expr=AssignExpr(
|
||||
location=loc,
|
||||
name=target,
|
||||
value=parsed_value,
|
||||
),
|
||||
)
|
||||
targets=[
|
||||
VariableExpr(
|
||||
location=Location.from_ast(node.target), name=target
|
||||
),
|
||||
],
|
||||
value=self.parse_expr(value),
|
||||
),
|
||||
)
|
||||
case _:
|
||||
print(f"Unsupported annotation: {ast.unparse(node)}")
|
||||
return statements
|
||||
|
||||
def parse_assign(self, node: ast.Assign) -> AssignStmt:
|
||||
targets: list[Expr] = []
|
||||
for target in node.targets:
|
||||
targets.append(self.parse_expr(target))
|
||||
value: Expr = self.parse_expr(node.value)
|
||||
return AssignStmt(
|
||||
location=Location.from_ast(node),
|
||||
targets=targets,
|
||||
value=value,
|
||||
)
|
||||
|
||||
def parse_function(self, node: ast.FunctionDef) -> Function:
|
||||
loc: Location = Location.from_ast(node)
|
||||
match node:
|
||||
@@ -228,4 +257,87 @@ class PythonParser:
|
||||
raise UnsupportedSyntaxError(column)
|
||||
|
||||
def parse_expr(self, node: ast.expr) -> Expr:
|
||||
raise NotImplementedError()
|
||||
match node:
|
||||
case ast.BoolOp():
|
||||
return self.parse_bool_op(node)
|
||||
|
||||
case ast.BinOp(left=left, op=op, right=right):
|
||||
return BinaryExpr(
|
||||
left=self.parse_expr(left),
|
||||
operator=op,
|
||||
right=self.parse_expr(right),
|
||||
)
|
||||
|
||||
case ast.UnaryOp(op=op, operand=right):
|
||||
return UnaryExpr(
|
||||
operator=op,
|
||||
right=self.parse_expr(right),
|
||||
)
|
||||
|
||||
case ast.Compare():
|
||||
return self.parse_compare(node)
|
||||
|
||||
case ast.Call():
|
||||
return self.parse_call(node)
|
||||
|
||||
case ast.Constant(value=value):
|
||||
return LiteralExpr(value=value)
|
||||
|
||||
case ast.Attribute(value=object, attr=name):
|
||||
return GetExpr(
|
||||
object=self.parse_expr(object),
|
||||
name=name,
|
||||
)
|
||||
|
||||
case ast.Name(id=name):
|
||||
return VariableExpr(name=name)
|
||||
|
||||
case _:
|
||||
raise UnsupportedSyntaxError(node)
|
||||
|
||||
def parse_bool_op(self, node: ast.BoolOp) -> LogicalExpr:
|
||||
op: ast.boolop = node.op
|
||||
values: list[ast.expr] = node.values
|
||||
expr: LogicalExpr = LogicalExpr(
|
||||
left=self.parse_expr(values[0]),
|
||||
operator=op,
|
||||
right=self.parse_expr(values[1]),
|
||||
)
|
||||
for value in values[2:]:
|
||||
expr = LogicalExpr(
|
||||
left=expr,
|
||||
operator=op,
|
||||
right=self.parse_expr(value),
|
||||
)
|
||||
return expr
|
||||
|
||||
def parse_compare(self, node: ast.Compare) -> Expr:
|
||||
ops: list[ast.cmpop] = node.ops
|
||||
rights: list[Expr] = [self.parse_expr(expr) for expr in node.comparators]
|
||||
expr: Expr = CompareExpr(
|
||||
left=self.parse_expr(node.left),
|
||||
operator=ops[0],
|
||||
right=rights[0],
|
||||
)
|
||||
for i, right in enumerate(rights[1:]):
|
||||
expr = LogicalExpr(
|
||||
left=expr,
|
||||
operator=ast.And(),
|
||||
right=CompareExpr(
|
||||
left=rights[i],
|
||||
operator=ops[i],
|
||||
right=right,
|
||||
),
|
||||
)
|
||||
return expr
|
||||
|
||||
def parse_call(self, node: ast.Call) -> CallExpr:
|
||||
return CallExpr(
|
||||
callee=self.parse_expr(node.func),
|
||||
arguments=[self.parse_expr(arg) for arg in node.args],
|
||||
keywords={
|
||||
arg.arg: self.parse_expr(arg.value)
|
||||
for arg in node.keywords
|
||||
if arg.arg is not None # Should always be True, type checker happy
|
||||
},
|
||||
)
|
||||
Reference in new issue
Block a user