Compare commits

..
4 Commits
11 changed files with 181 additions and 17 deletions

No files matched your search

+4
View File
@@ -69,6 +69,10 @@ class AssignStmt:
value: Expr value: Expr
class ReturnStmt:
value: Optional[Expr]
###< ###<
+5
View File
@@ -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():
+11
View File
@@ -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 #
############### ###############
+83 -2
View File
@@ -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:
+1
View File
@@ -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
View File
@@ -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
+2
View File
@@ -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
View File
@@ -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)
+6 -1
View File
@@ -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
+7
View File
@@ -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
+4
View File
@@ -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)