Compare commits

...
2 Commits
Author SHA1 Message Date
HEL 0b3f33d7fe feat(parser): parse python expressions 2026-05-25 23:17:52 +02:00
HEL 8a9b4f3989 feat(parser): parse assignments 2026-05-25 22:43:38 +02:00
5 changed files with 201 additions and 38 deletions

No files matched your search

+11 -5
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+3 -1
View File
@@ -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
View File
@@ -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
},
)