Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
fd5399f50a
|
||
|
|
8906ac3db8
|
||
|
|
022aebf55b
|
||
|
|
5dc6903425
|
No files matched your search
@@ -69,6 +69,10 @@ class AssignStmt:
|
|||||||
value: Expr
|
value: Expr
|
||||||
|
|
||||||
|
|
||||||
|
class ReturnStmt:
|
||||||
|
value: Optional[Expr]
|
||||||
|
|
||||||
|
|
||||||
###<
|
###<
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -449,6 +449,11 @@ class PythonAstPrinter(
|
|||||||
with self._child_level(single=True):
|
with self._child_level(single=True):
|
||||||
stmt.value.accept(self)
|
stmt.value.accept(self)
|
||||||
|
|
||||||
|
def visit_return_stmt(self, stmt: p.ReturnStmt) -> None:
|
||||||
|
self._write_line("ReturnStmt")
|
||||||
|
with self._child_level():
|
||||||
|
self._write_optional_child("value", stmt.value, last=True)
|
||||||
|
|
||||||
def visit_binary_expr(self, expr: p.BinaryExpr) -> None:
|
def visit_binary_expr(self, expr: p.BinaryExpr) -> None:
|
||||||
self._write_line("BinaryExpr")
|
self._write_line("BinaryExpr")
|
||||||
with self._child_level():
|
with self._child_level():
|
||||||
|
|||||||
@@ -100,6 +100,9 @@ class Stmt(ABC):
|
|||||||
@abstractmethod
|
@abstractmethod
|
||||||
def visit_assign_stmt(self, stmt: AssignStmt) -> T: ...
|
def visit_assign_stmt(self, stmt: AssignStmt) -> T: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def visit_return_stmt(self, stmt: ReturnStmt) -> T: ...
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class ExpressionStmt(Stmt):
|
class ExpressionStmt(Stmt):
|
||||||
@@ -150,6 +153,14 @@ class AssignStmt(Stmt):
|
|||||||
return visitor.visit_assign_stmt(self)
|
return visitor.visit_assign_stmt(self)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ReturnStmt(Stmt):
|
||||||
|
value: Optional[Expr]
|
||||||
|
|
||||||
|
def accept(self, visitor: Stmt.Visitor[T]) -> T:
|
||||||
|
return visitor.visit_return_stmt(self)
|
||||||
|
|
||||||
|
|
||||||
###############
|
###############
|
||||||
# Expressions #
|
# Expressions #
|
||||||
###############
|
###############
|
||||||
|
|||||||
@@ -8,13 +8,17 @@ from midas.ast.location import Location
|
|||||||
from midas.checker.diagnostic import Diagnostic, DiagnosticType
|
from midas.checker.diagnostic import Diagnostic, DiagnosticType
|
||||||
from midas.checker.environment import Environment
|
from midas.checker.environment import Environment
|
||||||
from midas.checker.operators import OPERATOR_METHODS
|
from midas.checker.operators import OPERATOR_METHODS
|
||||||
from midas.checker.types import Type, UnknownType
|
from midas.checker.types import Function, Type, UnitType, UnknownType
|
||||||
from midas.lexer.midas import MidasLexer
|
from midas.lexer.midas import MidasLexer
|
||||||
from midas.lexer.token import Token
|
from midas.lexer.token import Token
|
||||||
from midas.parser.midas import MidasParser
|
from midas.parser.midas import MidasParser
|
||||||
from midas.resolver.midas import MidasResolver
|
from midas.resolver.midas import MidasResolver
|
||||||
|
|
||||||
|
|
||||||
|
class ReturnException(Exception):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class Checker(
|
class Checker(
|
||||||
p.Stmt.Visitor[None],
|
p.Stmt.Visitor[None],
|
||||||
p.Expr.Visitor[Type],
|
p.Expr.Visitor[Type],
|
||||||
@@ -63,6 +67,16 @@ class Checker(
|
|||||||
def evaluate(self, expr: p.Expr) -> Type:
|
def evaluate(self, expr: p.Expr) -> Type:
|
||||||
return expr.accept(self)
|
return expr.accept(self)
|
||||||
|
|
||||||
|
def evaluate_block(self, block: list[p.Stmt], env: Environment) -> None:
|
||||||
|
previous_env: Environment = self.env
|
||||||
|
self.env = env
|
||||||
|
for stmt in block:
|
||||||
|
try:
|
||||||
|
stmt.accept(self)
|
||||||
|
except ReturnException:
|
||||||
|
break
|
||||||
|
self.env = previous_env
|
||||||
|
|
||||||
def check(self, statements: list[p.Stmt]) -> list[Diagnostic]:
|
def check(self, statements: list[p.Stmt]) -> list[Diagnostic]:
|
||||||
self.diagnostics = []
|
self.diagnostics = []
|
||||||
for stmt in statements:
|
for stmt in statements:
|
||||||
@@ -105,7 +119,69 @@ class Checker(
|
|||||||
def visit_expression_stmt(self, stmt: p.ExpressionStmt) -> None:
|
def visit_expression_stmt(self, stmt: p.ExpressionStmt) -> None:
|
||||||
self.evaluate(stmt.expr)
|
self.evaluate(stmt.expr)
|
||||||
|
|
||||||
def visit_function(self, stmt: p.Function) -> None: ...
|
def visit_function(self, stmt: p.Function) -> None:
|
||||||
|
env: Environment = Environment(self.env)
|
||||||
|
pos_args: list[Function.Argument] = []
|
||||||
|
args: list[Function.Argument] = []
|
||||||
|
kw_args: list[Function.Argument] = []
|
||||||
|
|
||||||
|
def eval_arg_type(arg: p.Function.Argument) -> Type:
|
||||||
|
if arg.type is None:
|
||||||
|
return UnknownType()
|
||||||
|
return arg.type.accept(self)
|
||||||
|
|
||||||
|
for arg in stmt.posonlyargs:
|
||||||
|
pos_args.append(
|
||||||
|
Function.Argument(
|
||||||
|
name=arg.name,
|
||||||
|
type=eval_arg_type(arg),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
for arg in stmt.args:
|
||||||
|
args.append(
|
||||||
|
Function.Argument(
|
||||||
|
name=arg.name,
|
||||||
|
type=eval_arg_type(arg),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
for arg in stmt.kwonlyargs:
|
||||||
|
kw_args.append(
|
||||||
|
Function.Argument(
|
||||||
|
name=arg.name,
|
||||||
|
type=eval_arg_type(arg),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
for arg in pos_args + args + kw_args:
|
||||||
|
env.define(arg.name, arg.type)
|
||||||
|
|
||||||
|
self.evaluate_block(stmt.body, env)
|
||||||
|
inferred_return: Type = UnknownType()
|
||||||
|
if len(env.return_types) == 1:
|
||||||
|
inferred_return = list(env.return_types)[0]
|
||||||
|
elif len(env.return_types) > 1:
|
||||||
|
self.error(
|
||||||
|
stmt.location,
|
||||||
|
f"Mixed return types: {env.return_types}",
|
||||||
|
)
|
||||||
|
returns: Type = UnknownType()
|
||||||
|
if stmt.returns is not None:
|
||||||
|
returns = stmt.returns.accept(self)
|
||||||
|
if returns != inferred_return:
|
||||||
|
self.error(
|
||||||
|
stmt.returns.location,
|
||||||
|
f"Return type mismatch, annotated {returns} but returns {inferred_return}",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
returns = inferred_return
|
||||||
|
|
||||||
|
function: Function = Function(
|
||||||
|
pos_args=pos_args,
|
||||||
|
args=args,
|
||||||
|
kw_args=kw_args,
|
||||||
|
returns=returns,
|
||||||
|
)
|
||||||
|
self.env.define(stmt.name, function)
|
||||||
|
|
||||||
def visit_type_assign(self, stmt: p.TypeAssign) -> None:
|
def visit_type_assign(self, stmt: p.TypeAssign) -> None:
|
||||||
# TODO check not yet defined locally
|
# TODO check not yet defined locally
|
||||||
@@ -132,6 +208,11 @@ class Checker(
|
|||||||
f"Cannot assign {value} to {name} of type {var_type}",
|
f"Cannot assign {value} to {name} of type {var_type}",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def visit_return_stmt(self, stmt: p.ReturnStmt) -> None:
|
||||||
|
type: Type = stmt.value.accept(self) if stmt.value is not None else UnitType()
|
||||||
|
self.env.return_types.add(type)
|
||||||
|
raise ReturnException()
|
||||||
|
|
||||||
def visit_binary_expr(self, expr: p.BinaryExpr) -> Type:
|
def visit_binary_expr(self, expr: p.BinaryExpr) -> Type:
|
||||||
method: Optional[str] = OPERATOR_METHODS.get(expr.operator.__class__)
|
method: Optional[str] = OPERATOR_METHODS.get(expr.operator.__class__)
|
||||||
if method is None:
|
if method is None:
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ class Environment:
|
|||||||
def __init__(self, enclosing: Optional[Environment] = None) -> None:
|
def __init__(self, enclosing: Optional[Environment] = None) -> None:
|
||||||
self.enclosing: Optional[Environment] = enclosing
|
self.enclosing: Optional[Environment] = enclosing
|
||||||
self.values: dict[str, Type] = {}
|
self.values: dict[str, Type] = {}
|
||||||
|
self.return_types: set[Type] = set()
|
||||||
|
|
||||||
def define(self, name: str, value: Type):
|
def define(self, name: str, value: Type):
|
||||||
self.values[name] = value
|
self.values[name] = value
|
||||||
|
|||||||
+19
-1
@@ -19,4 +19,22 @@ class UnknownType:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
Type = BaseType | SimpleType | UnknownType
|
@dataclass(frozen=True, kw_only=True)
|
||||||
|
class UnitType:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, kw_only=True)
|
||||||
|
class Function:
|
||||||
|
pos_args: list[Argument]
|
||||||
|
args: list[Argument]
|
||||||
|
kw_args: list[Argument]
|
||||||
|
returns: Type
|
||||||
|
|
||||||
|
@dataclass(frozen=True, kw_only=True)
|
||||||
|
class Argument:
|
||||||
|
name: str
|
||||||
|
type: Type
|
||||||
|
|
||||||
|
|
||||||
|
Type = BaseType | SimpleType | UnknownType | UnitType | Function
|
||||||
@@ -153,6 +153,8 @@ class PythonHighlighter(
|
|||||||
|
|
||||||
def visit_assign_stmt(self, stmt: p.AssignStmt) -> None: ...
|
def visit_assign_stmt(self, stmt: p.AssignStmt) -> None: ...
|
||||||
|
|
||||||
|
def visit_return_stmt(self, stmt: p.ReturnStmt) -> None: ...
|
||||||
|
|
||||||
def visit_binary_expr(self, expr: p.BinaryExpr) -> None: ...
|
def visit_binary_expr(self, expr: p.BinaryExpr) -> None: ...
|
||||||
|
|
||||||
def visit_compare_expr(self, expr: p.CompareExpr) -> None: ...
|
def visit_compare_expr(self, expr: p.CompareExpr) -> None: ...
|
||||||
|
|||||||
+39
-13
@@ -9,7 +9,7 @@ import click
|
|||||||
import midas.ast.midas as m
|
import midas.ast.midas as m
|
||||||
import midas.ast.python as p
|
import midas.ast.python as p
|
||||||
from midas.ast.location import Location
|
from midas.ast.location import Location
|
||||||
from midas.ast.printer import PythonAstPrinter
|
from midas.ast.printer import MidasAstPrinter, PythonAstPrinter
|
||||||
from midas.checker.checker import Checker
|
from midas.checker.checker import Checker
|
||||||
from midas.checker.diagnostic import Diagnostic
|
from midas.checker.diagnostic import Diagnostic
|
||||||
from midas.cli.highlighter import Highlighter, MidasHighlighter, PythonHighlighter
|
from midas.cli.highlighter import Highlighter, MidasHighlighter, PythonHighlighter
|
||||||
@@ -46,26 +46,52 @@ def utils():
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def dump_python_ast(tree: ast.Module) -> str:
|
||||||
|
parser = PythonParser()
|
||||||
|
stmts: list[p.Stmt] = parser.parse_module(tree)
|
||||||
|
printer = PythonAstPrinter()
|
||||||
|
dump: str = ""
|
||||||
|
for stmt in stmts:
|
||||||
|
dump += printer.print(stmt)
|
||||||
|
dump += "\n"
|
||||||
|
return dump
|
||||||
|
|
||||||
|
|
||||||
|
def dump_midas_ast(source: str, filename: str) -> str:
|
||||||
|
lexer = MidasLexer(source, file=filename)
|
||||||
|
tokens: list[Token] = lexer.process()
|
||||||
|
parser = MidasParser(tokens)
|
||||||
|
stmts: list[m.Stmt] = parser.parse()
|
||||||
|
if len(parser.errors) != 0:
|
||||||
|
for err in parser.errors:
|
||||||
|
print(err.get_report())
|
||||||
|
raise RuntimeError("A parsing error occurred")
|
||||||
|
printer = MidasAstPrinter()
|
||||||
|
dump: str = ""
|
||||||
|
for stmt in stmts:
|
||||||
|
dump += printer.print(stmt)
|
||||||
|
dump += "\n"
|
||||||
|
return dump
|
||||||
|
|
||||||
|
|
||||||
@utils.command()
|
@utils.command()
|
||||||
@click.option("-o", "--output", type=click.File("w"))
|
@click.option("-o", "--output", type=click.File("w"))
|
||||||
@click.option("-p", "--parse", is_flag=True)
|
@click.option("-p", "--parse", is_flag=True)
|
||||||
@click.argument("file", type=click.File("r"))
|
@click.argument("file", type=click.File("r"))
|
||||||
def dump_ast(output: Optional[TextIO], parse: bool, file: TextIO):
|
def dump_ast(output: Optional[TextIO], parse: bool, file: TextIO):
|
||||||
source: str = file.read()
|
source: str = file.read()
|
||||||
tree: ast.Module = ast.parse(source, filename=file.name)
|
|
||||||
dump: str
|
dump: str
|
||||||
|
if file.name.endswith(".py"):
|
||||||
if parse:
|
tree: ast.Module = ast.parse(source, filename=file.name)
|
||||||
parser = PythonParser()
|
if parse:
|
||||||
stmts: list[p.Stmt] = parser.parse_module(tree)
|
dump = dump_python_ast(tree)
|
||||||
printer = PythonAstPrinter()
|
else:
|
||||||
dump = ""
|
dump = ast.dump(tree, indent=4)
|
||||||
for stmt in stmts:
|
elif file.name.endswith(".midas"):
|
||||||
dump += printer.print(stmt)
|
dump = dump_midas_ast(source, file.name)
|
||||||
dump += "\n"
|
|
||||||
|
|
||||||
else:
|
else:
|
||||||
dump = ast.dump(tree, indent=4)
|
raise ValueError("Unsupported file type")
|
||||||
|
|
||||||
if output is None:
|
if output is None:
|
||||||
click.echo(dump)
|
click.echo(dump)
|
||||||
|
|||||||
@@ -319,8 +319,13 @@ class MidasParser(Parser):
|
|||||||
"""
|
"""
|
||||||
self.consume(TokenType.LEFT_BRACE, "Expected '{' to start type body")
|
self.consume(TokenType.LEFT_BRACE, "Expected '{' to start type body")
|
||||||
properties: list[PropertyStmt] = []
|
properties: list[PropertyStmt] = []
|
||||||
|
names: set[str] = set()
|
||||||
while not self.check(TokenType.RIGHT_BRACE) and not self.is_at_end():
|
while not self.check(TokenType.RIGHT_BRACE) and not self.is_at_end():
|
||||||
properties.append(self.property_stmt())
|
prop: PropertyStmt = self.property_stmt()
|
||||||
|
if prop.name.lexeme in names:
|
||||||
|
raise self.error(prop.name, "Duplicate property")
|
||||||
|
names.add(prop.name.lexeme)
|
||||||
|
properties.append(prop)
|
||||||
self.consume(TokenType.RIGHT_BRACE, "Unclosed type body")
|
self.consume(TokenType.RIGHT_BRACE, "Unclosed type body")
|
||||||
return properties
|
return properties
|
||||||
|
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ from midas.ast.python import (
|
|||||||
LiteralExpr,
|
LiteralExpr,
|
||||||
LogicalExpr,
|
LogicalExpr,
|
||||||
MidasType,
|
MidasType,
|
||||||
|
ReturnStmt,
|
||||||
Stmt,
|
Stmt,
|
||||||
TypeAssign,
|
TypeAssign,
|
||||||
UnaryExpr,
|
UnaryExpr,
|
||||||
@@ -70,6 +71,12 @@ class PythonParser:
|
|||||||
expr=self.parse_expr(expr),
|
expr=self.parse_expr(expr),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
case ast.Return(value=value):
|
||||||
|
return ReturnStmt(
|
||||||
|
location=location,
|
||||||
|
value=self.parse_expr(value) if value is not None else None,
|
||||||
|
)
|
||||||
|
|
||||||
case _:
|
case _:
|
||||||
print(f"Unsupported statement: {ast.unparse(node)}")
|
print(f"Unsupported statement: {ast.unparse(node)}")
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -72,6 +72,10 @@ class Resolver(p.Stmt.Visitor[None], p.Expr.Visitor[None]):
|
|||||||
case _:
|
case _:
|
||||||
raise Exception(f"Unsupported assignment to {target}")
|
raise Exception(f"Unsupported assignment to {target}")
|
||||||
|
|
||||||
|
def visit_return_stmt(self, stmt: p.ReturnStmt) -> None:
|
||||||
|
if stmt.value is not None:
|
||||||
|
self.resolve(stmt.value)
|
||||||
|
|
||||||
def visit_binary_expr(self, expr: p.BinaryExpr) -> None:
|
def visit_binary_expr(self, expr: p.BinaryExpr) -> None:
|
||||||
self.resolve(expr.left)
|
self.resolve(expr.left)
|
||||||
self.resolve(expr.right)
|
self.resolve(expr.right)
|
||||||
|
|||||||
Reference in new issue
Block a user