Compare commits

...
5 Commits
Author SHA1 Message Date
HEL 4d23e8840e feat(parser): adapt AST printer with new nodes 2026-05-25 22:06:18 +02:00
HEL c64d626d1c refactor(parser): remove inheritance from NodeVisitor
remove the parent NodeVisitor class from PythonParser and implement all custom recursive methods instead
2026-05-25 21:42:04 +02:00
HEL ecab1b74a4 feat(parser): add Python AST nodes 2026-05-25 21:39:20 +02:00
HEL 0bbdf04621 feat(parser): generate python AST classes
use the generation script to create Python AST node classes, also distinguish between Midas type annotation nodes and statements
2026-05-25 20:53:36 +02:00
HEL 939e5af4ce refactor(parser): improve AST class generator
make the generation script more flexible
2026-05-25 20:38:38 +02:00
7 changed files with 644 additions and 180 deletions

No files matched your search

+78 -68
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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()