Compare commits

...
3 Commits
Author SHA1 Message Date
HEL 1b078b832c chore: add some operations in the example 2026-05-28 18:32:35 +02:00
HEL 7515716864 feat(checker): add diagnostics 2026-05-28 18:32:35 +02:00
HEL 218b0c5b78 fix(parser): add location in all AST nodes 2026-05-28 18:32:34 +02:00
10 changed files with 150 additions and 53 deletions

No files matched your search

@@ -2,3 +2,10 @@ a: int = 3
b: int = 4 b: int = 4
c = a + b # -> int c = a + b # -> int
c = "invalid" # -> can't assign str to int variable
d = True
e = d + d
f: float = a
+1 -1
View File
@@ -11,7 +11,7 @@ SECTION_TEMPLATE = """{banner}
@dataclass(frozen=True, kw_only=True) @dataclass(frozen=True, kw_only=True)
class {base}(ABC): class {base}(ABC):
location: Optional[Location] = None location: Location
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
+2 -2
View File
@@ -21,7 +21,7 @@ T = TypeVar("T")
@dataclass(frozen=True, kw_only=True) @dataclass(frozen=True, kw_only=True)
class Stmt(ABC): class Stmt(ABC):
location: Optional[Location] = None location: Location
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
@@ -114,7 +114,7 @@ class PredicateStmt(Stmt):
@dataclass(frozen=True, kw_only=True) @dataclass(frozen=True, kw_only=True)
class Expr(ABC): class Expr(ABC):
location: Optional[Location] = None location: Location
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
+3 -3
View File
@@ -21,7 +21,7 @@ T = TypeVar("T")
@dataclass(frozen=True, kw_only=True) @dataclass(frozen=True, kw_only=True)
class MidasType(ABC): class MidasType(ABC):
location: Optional[Location] = None location: Location
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
@@ -82,7 +82,7 @@ class FrameType(MidasType):
@dataclass(frozen=True, kw_only=True) @dataclass(frozen=True, kw_only=True)
class Stmt(ABC): class Stmt(ABC):
location: Optional[Location] = None location: Location
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
@@ -157,7 +157,7 @@ class AssignStmt(Stmt):
@dataclass(frozen=True, kw_only=True) @dataclass(frozen=True, kw_only=True)
class Expr(ABC): class Expr(ABC):
location: Optional[Location] = None location: Location
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
+54 -8
View File
@@ -4,9 +4,11 @@ from typing import Optional
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.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 BaseType, Type, UnknownType from midas.checker.types import Type, 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
@@ -18,22 +20,56 @@ class Checker(
p.Expr.Visitor[Type], p.Expr.Visitor[Type],
p.MidasType.Visitor[Type], p.MidasType.Visitor[Type],
): ):
def __init__(self, locals: dict[p.Expr, int], base_dir: Path): def __init__(self, locals: dict[p.Expr, int], file_path: Path):
self.logger: logging.Logger = logging.getLogger("Checker") self.logger: logging.Logger = logging.getLogger("Checker")
self.base_dir: Path = base_dir self.file_path: Path = file_path
self.ctx: MidasResolver = MidasResolver() self.ctx: MidasResolver = MidasResolver()
self.global_env: Environment = Environment() self.global_env: Environment = Environment()
self.env: Environment = self.global_env self.env: Environment = self.global_env
self.locals: dict[p.Expr, int] = locals self.locals: dict[p.Expr, int] = locals
self.diagnostics: list[Diagnostic] = []
def diagnostic(self, type: DiagnosticType, location: Location, message: str):
self.diagnostics.append(
Diagnostic(
file_path=self.file_path,
location=location,
type=type,
message=message,
)
)
def error(self, location: Location, message: str):
self.diagnostic(
type=DiagnosticType.ERROR,
location=location,
message=message,
)
def warning(self, location: Location, message: str):
self.diagnostic(
type=DiagnosticType.WARNING,
location=location,
message=message,
)
def info(self, location: Location, message: str):
self.diagnostic(
type=DiagnosticType.INFO,
location=location,
message=message,
)
def evaluate(self, expr: p.Expr) -> Type: def evaluate(self, expr: p.Expr) -> Type:
return expr.accept(self) return expr.accept(self)
def check(self, statements: list[p.Stmt]) -> None: def check(self, statements: list[p.Stmt]) -> list[Diagnostic]:
self.diagnostics = []
for stmt in statements: for stmt in statements:
stmt.accept(self) stmt.accept(self)
self.logger.debug(f"Final environment: {self.env.flat_dict()}") self.logger.debug(f"Final environment: {self.env.flat_dict()}")
return self.diagnostics
def look_up_variable(self, name: str, expr: p.Expr) -> Optional[Type]: def look_up_variable(self, name: str, expr: p.Expr) -> Optional[Type]:
distance: Optional[int] = self.locals.get(expr) distance: Optional[int] = self.locals.get(expr)
@@ -57,7 +93,7 @@ class Checker(
def import_midas(self, path: Path) -> None: def import_midas(self, path: Path) -> None:
self.logger.debug(f"Importing type definitions from {path}") self.logger.debug(f"Importing type definitions from {path}")
path = (self.base_dir / path).resolve() path = (self.file_path.parent / path).resolve()
lexer: MidasLexer = MidasLexer(path.read_text()) lexer: MidasLexer = MidasLexer(path.read_text())
tokens: list[Token] = lexer.process() tokens: list[Token] = lexer.process()
parser: MidasParser = MidasParser(tokens) parser: MidasParser = MidasParser(tokens)
@@ -81,6 +117,7 @@ class Checker(
for target in stmt.targets: for target in stmt.targets:
if not isinstance(target, p.VariableExpr): if not isinstance(target, p.VariableExpr):
self.logger.warning(f"Unsupported assignment to {target}") self.logger.warning(f"Unsupported assignment to {target}")
self.warning(target.location, f"Unsupported assignment to {target}")
continue continue
name: str = target.name name: str = target.name
var_type: Optional[Type] = self.look_up_variable(name, target) var_type: Optional[Type] = self.look_up_variable(name, target)
@@ -90,19 +127,27 @@ class Checker(
else: else:
# TODO: implement real comparison method # TODO: implement real comparison method
if var_type != value: if var_type != value:
raise ValueError( self.error(
f"Cannot assign {value} to {name} of type {var_type}" stmt.location,
f"Cannot assign {value} to {name} of type {var_type}",
) )
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:
self.logger.warning(f"Unsupported operator {expr.operator}") self.logger.warning(f"Unsupported operator {expr.operator}")
self.warning(expr.location, f"Unsupported operator {expr.operator}")
return UnknownType() return UnknownType()
left: Type = self.evaluate(expr.left) left: Type = self.evaluate(expr.left)
right: Type = self.evaluate(expr.right) right: Type = self.evaluate(expr.right)
result: Type = self.ctx.get_operation_result(left, method, right) result: Optional[Type] = self.ctx.get_operation_result(left, method, right)
if result is None:
self.error(
expr.location,
f"Undefined operation {method} between {left} and {right}",
)
return UnknownType()
return result return result
def visit_compare_expr(self, expr: p.CompareExpr) -> Type: ... def visit_compare_expr(self, expr: p.CompareExpr) -> Type: ...
@@ -133,6 +178,7 @@ class Checker(
case str(): case str():
return self.ctx.get_type("str") return self.ctx.get_type("str")
case _: case _:
self.warning(expr.location, f"Unknown literal {expr}")
return UnknownType() return UnknownType()
def visit_variable_expr(self, expr: p.VariableExpr) -> Type: def visit_variable_expr(self, expr: p.VariableExpr) -> Type:
+33
View File
@@ -0,0 +1,33 @@
from dataclasses import dataclass
from enum import StrEnum
from pathlib import Path
from typing import Optional
from midas.ast.location import Location
class DiagnosticType(StrEnum):
ERROR = "Error"
WARNING = "Warning"
INFO = "Info"
@dataclass(frozen=True)
class Diagnostic:
file_path: Path
location: Location
type: DiagnosticType
message: str
def __str__(self) -> str:
start_loc: str = f"L{self.location.lineno}:{self.location.col_offset+1}"
end_loc: Optional[str] = ""
if (
self.location.end_lineno is not None
and self.location.end_col_offset is not None
):
end_loc = f"L{self.location.end_lineno}:{self.location.end_col_offset+1}"
loc: str = (
f"at {start_loc}" if end_loc is None else f"from {start_loc} to {end_loc}"
)
return f"{self.type} in {self.file_path} {loc}: {self.message}"
+5 -2
View File
@@ -11,6 +11,7 @@ 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 PythonAstPrinter
from midas.checker.checker import Checker from midas.checker.checker import Checker
from midas.checker.diagnostic import Diagnostic
from midas.cli.highlighter import Highlighter, MidasHighlighter, PythonHighlighter from midas.cli.highlighter import Highlighter, MidasHighlighter, PythonHighlighter
from midas.lexer.midas import MidasLexer from midas.lexer.midas import MidasLexer
from midas.lexer.token import Token, TokenType from midas.lexer.token import Token, TokenType
@@ -34,8 +35,10 @@ def compile(file: TextIO):
stmts: list[p.Stmt] = parser.parse_module(tree) stmts: list[p.Stmt] = parser.parse_module(tree)
resolver = Resolver() resolver = Resolver()
resolver.resolve(*stmts) resolver.resolve(*stmts)
checker = Checker(resolver.locals, base_dir=Path(file.name).resolve().parent) checker = Checker(resolver.locals, file_path=Path(file.name).resolve())
checker.check(stmts) diagnostics: list[Diagnostic] = checker.check(stmts)
for diagnostic in diagnostics:
print(diagnostic)
@midas.group() @midas.group()
+6 -18
View File
@@ -205,9 +205,7 @@ class MidasParser(Parser):
while self.match(TokenType.AND): while self.match(TokenType.AND):
operator: Token = self.previous() operator: Token = self.previous()
right: Expr = self.equality() right: Expr = self.equality()
location: Optional[Location] = None location: Location = Location.span(expr.location, right.location)
if expr.location and right.location:
location = Location.span(expr.location, right.location)
expr = LogicalExpr( expr = LogicalExpr(
location=location, left=expr, operator=operator, right=right location=location, left=expr, operator=operator, right=right
) )
@@ -223,9 +221,7 @@ class MidasParser(Parser):
while self.match(TokenType.BANG_EQUAL, TokenType.EQUAL_EQUAL): while self.match(TokenType.BANG_EQUAL, TokenType.EQUAL_EQUAL):
operator: Token = self.previous() operator: Token = self.previous()
right: Expr = self.comparison() right: Expr = self.comparison()
location: Optional[Location] = None location: Location = Location.span(expr.location, right.location)
if expr.location and right.location:
location = Location.span(expr.location, right.location)
expr = BinaryExpr( expr = BinaryExpr(
location=location, left=expr, operator=operator, right=right location=location, left=expr, operator=operator, right=right
) )
@@ -246,9 +242,7 @@ class MidasParser(Parser):
): ):
operator: Token = self.previous() operator: Token = self.previous()
right: Expr = self.unary() right: Expr = self.unary()
location: Optional[Location] = None location: Location = Location.span(expr.location, right.location)
if expr.location and right.location:
location = Location.span(expr.location, right.location)
expr = BinaryExpr( expr = BinaryExpr(
location=location, left=expr, operator=operator, right=right location=location, left=expr, operator=operator, right=right
) )
@@ -263,9 +257,7 @@ class MidasParser(Parser):
if self.match(TokenType.MINUS): if self.match(TokenType.MINUS):
operator: Token = self.previous() operator: Token = self.previous()
right: Expr = self.unary() right: Expr = self.unary()
location: Optional[Location] = None location: Location = Location.span(operator.get_location(), right.location)
if right.location:
location = Location.span(operator.get_location(), right.location)
return UnaryExpr(location=location, operator=operator, right=right) return UnaryExpr(location=location, operator=operator, right=right)
return self.reference() return self.reference()
@@ -280,9 +272,7 @@ class MidasParser(Parser):
name: Token = self.consume( name: Token = self.consume(
TokenType.IDENTIFIER, "Expected property name after '.'" TokenType.IDENTIFIER, "Expected property name after '.'"
) )
location: Optional[Location] = None location: Location = Location.span(expr.location, name.get_location())
if expr.location:
location = Location.span(expr.location, name.get_location())
expr = GetExpr(location=location, expr=expr, name=name) expr = GetExpr(location=location, expr=expr, name=name)
return expr return expr
@@ -370,9 +360,7 @@ class MidasParser(Parser):
while not self.is_at_end() and not self.check(TokenType.RIGHT_BRACE): while not self.is_at_end() and not self.check(TokenType.RIGHT_BRACE):
operations.append(self.op_declaration()) operations.append(self.op_declaration())
self.consume(TokenType.RIGHT_BRACE, "Unclosed extend body") self.consume(TokenType.RIGHT_BRACE, "Unclosed extend body")
location: Optional[Location] = None location: Location = keyword.location_to(self.previous())
if type.location:
location = keyword.location_to(self.previous())
return ExtendStmt(location=location, type=type, operations=operations) return ExtendStmt(location=location, type=type, operations=operations)
def op_declaration(self) -> OpStmt: def op_declaration(self) -> OpStmt:
+36 -14
View File
@@ -53,6 +53,7 @@ class PythonParser:
return statements return statements
def parse_stmt(self, node: ast.stmt) -> None | Stmt | list[Stmt]: def parse_stmt(self, node: ast.stmt) -> None | Stmt | list[Stmt]:
location: Location = Location.from_ast(node)
match node: match node:
case ast.AnnAssign(): case ast.AnnAssign():
return self.parse_annotation_assign(node) return self.parse_annotation_assign(node)
@@ -64,7 +65,10 @@ class PythonParser:
return self.parse_function(node) return self.parse_function(node)
case ast.Expr(value=expr): case ast.Expr(value=expr):
return ExpressionStmt(expr=self.parse_expr(expr)) return ExpressionStmt(
location=location,
expr=self.parse_expr(expr),
)
case _: case _:
print(f"Unsupported statement: {ast.unparse(node)}") print(f"Unsupported statement: {ast.unparse(node)}")
@@ -266,12 +270,14 @@ class PythonParser:
raise UnsupportedSyntaxError(column) raise UnsupportedSyntaxError(column)
def parse_expr(self, node: ast.expr) -> Expr: def parse_expr(self, node: ast.expr) -> Expr:
location: Location = Location.from_ast(node)
match node: match node:
case ast.BoolOp(): case ast.BoolOp():
return self.parse_bool_op(node) return self.parse_bool_op(node)
case ast.BinOp(left=left, op=op, right=right): case ast.BinOp(left=left, op=op, right=right):
return BinaryExpr( return BinaryExpr(
location=location,
left=self.parse_expr(left), left=self.parse_expr(left),
operator=op, operator=op,
right=self.parse_expr(right), right=self.parse_expr(right),
@@ -279,6 +285,7 @@ class PythonParser:
case ast.UnaryOp(op=op, operand=right): case ast.UnaryOp(op=op, operand=right):
return UnaryExpr( return UnaryExpr(
location=location,
operator=op, operator=op,
right=self.parse_expr(right), right=self.parse_expr(right),
) )
@@ -290,58 +297,73 @@ class PythonParser:
return self.parse_call(node) return self.parse_call(node)
case ast.Constant(value=value): case ast.Constant(value=value):
return LiteralExpr(value=value) return LiteralExpr(location=location, value=value)
case ast.Attribute(value=object, attr=name): case ast.Attribute(value=object, attr=name):
return GetExpr( return GetExpr(
location=location,
object=self.parse_expr(object), object=self.parse_expr(object),
name=name, name=name,
) )
case ast.Name(id=name): case ast.Name(id=name):
return VariableExpr(name=name) return VariableExpr(location=location, name=name)
case _: case _:
raise UnsupportedSyntaxError(node) raise UnsupportedSyntaxError(node)
def parse_bool_op(self, node: ast.BoolOp) -> LogicalExpr: def parse_bool_op(self, node: ast.BoolOp) -> LogicalExpr:
op: ast.boolop = node.op op: ast.boolop = node.op
values: list[ast.expr] = node.values rights: list[Expr] = [self.parse_expr(expr) for expr in node.values]
expr: LogicalExpr = LogicalExpr( expr: LogicalExpr = LogicalExpr(
left=self.parse_expr(values[0]), location=Location.span(
rights[0].location,
rights[1].location,
),
left=rights[0],
operator=op, operator=op,
right=self.parse_expr(values[1]), right=rights[1],
) )
for value in values[2:]: for right in rights[2:]:
expr = LogicalExpr( expr = LogicalExpr(
location=Location.span(expr.location, right.location),
left=expr, left=expr,
operator=op, operator=op,
right=self.parse_expr(value), right=right,
) )
return expr return expr
def parse_compare(self, node: ast.Compare) -> Expr: def parse_compare(self, node: ast.Compare) -> Expr:
ops: list[ast.cmpop] = node.ops ops: list[ast.cmpop] = node.ops
left: Expr = self.parse_expr(node.left)
rights: list[Expr] = [self.parse_expr(expr) for expr in node.comparators] rights: list[Expr] = [self.parse_expr(expr) for expr in node.comparators]
expr: Expr = CompareExpr( expr: Expr = CompareExpr(
left=self.parse_expr(node.left), location=Location.span(
left.location,
rights[0].location,
),
left=left,
operator=ops[0], operator=ops[0],
right=rights[0], right=rights[0],
) )
for i, right in enumerate(rights[1:]): for i, right in enumerate(rights[1:]):
comparison = CompareExpr(
location=Location.span(rights[i].location, right.location),
left=rights[i],
operator=ops[i],
right=right,
)
expr = LogicalExpr( expr = LogicalExpr(
location=Location.span(expr.location, comparison.location),
left=expr, left=expr,
operator=ast.And(), operator=ast.And(),
right=CompareExpr( right=comparison,
left=rights[i],
operator=ops[i],
right=right,
),
) )
return expr return expr
def parse_call(self, node: ast.Call) -> CallExpr: def parse_call(self, node: ast.Call) -> CallExpr:
return CallExpr( return CallExpr(
location=Location.from_ast(node),
callee=self.parse_expr(node.func), callee=self.parse_expr(node.func),
arguments=[self.parse_expr(arg) for arg in node.args], arguments=[self.parse_expr(arg) for arg in node.args],
keywords={ keywords={
+3 -5
View File
@@ -17,13 +17,11 @@ class MidasResolver(m.Stmt.Visitor[None], m.Expr.Visitor[Type]):
raise NameError(f"Undefined type {name}") raise NameError(f"Undefined type {name}")
return type return type
def get_operation_result(self, left: Type, operator: str, right: Type) -> Type: def get_operation_result(
self, left: Type, operator: str, right: Type
) -> Optional[Type]:
operation: tuple[Type, str, Type] = (left, operator, right) operation: tuple[Type, str, Type] = (left, operator, right)
result: Optional[Type] = self._operations.get(operation) result: Optional[Type] = self._operations.get(operation)
if result is None:
raise ValueError(
f"Undefined operation {operator} between {left} and {right}"
)
return result return result
def _define_builtin(self): def _define_builtin(self):