Compare commits

...
12 Commits
Author SHA1 Message Date
HEL 4dd76e30cc feat(checker): add diagnostics 2026-05-28 18:19:02 +02:00
HEL 9f31ab623e fix(parser): add location in all AST nodes 2026-05-28 18:18:16 +02:00
HEL 928901ef9c fix(checker): get literal types from context 2026-05-28 18:16:35 +02:00
HEL 4b62c78874 feat(cli): integrate checker in compile command 2026-05-28 17:35:38 +02:00
HEL f882eebaf5 feat(checker): add basic checker
still very basic but lays out the structure and help methods
2026-05-28 17:35:00 +02:00
HEL a872938405 feat(checker): add midas context resolver
this is still very basic and only handle a few expressions
notably, it doesn't support generics, option types, conditions, predicates nor complex types
2026-05-28 17:33:16 +02:00
HEL 146be72fd7 chore: add simple operation and type examples 2026-05-28 17:31:12 +02:00
HEL 6de54e1da1 feat(checker): add python scope resolver
adapted from Pebble
2026-05-28 17:30:16 +02:00
HEL c82b41a4df feat(checker): add environment manager
adapted from Pebble
2026-05-28 17:29:37 +02:00
HEL 8304760fe0 fix(parser): add function body and all_args property 2026-05-28 15:26:53 +02:00
HEL 6bf91db757 feat(checker): create basic type and operation structs 2026-05-28 15:25:48 +02:00
HEL 3f6b650a4b docs: create architecture diagram 2026-05-28 15:25:12 +02:00
18 changed files with 824 additions and 40 deletions

No files matched your search

+150
View File
@@ -0,0 +1,150 @@
#import "@preview/cetz:0.5.2": canvas, draw
#let diagram-only = false
#set document(
title: [Midas Architecture],
//author: "Louis Heredero",
)
#set text(
font: "Source Sans 3",
)
#let diagram = canvas({
let framed = draw.content.with(
padding: (x: .8em, y: 1em),
frame: "rect",
stroke: black,
)
let arrow = draw.line.with(mark: (end: ">", fill: black))
framed(
(0, 0),
name: "python-parser",
)[Python parser]
draw.content(
(rel: (0, 1), to: "python-parser.north"),
padding: 5pt,
anchor: "south",
name: "source-py",
)[_`source.py`_]
arrow("source-py", "python-parser")
framed(
(rel: (3, 0), to: "python-parser.east"),
anchor: "west",
name: "custom-parser",
align(center)[Custom python\ parser],
)
arrow("python-parser", "custom-parser", name: "arrow-python-ast")
draw.content(
"arrow-python-ast",
anchor: "south",
padding: 5pt,
)[`ast.Module`]
framed(
(rel: (-3, -2), to: "custom-parser.south"),
anchor: "east",
name: "python-resolver",
)[Python Resolver]
arrow(
"custom-parser",
((), "|-", "python-resolver.east"),
"python-resolver",
name: "arrow-python-custom-ast",
)
draw.content(
(rel: (1.5, 0), to: "arrow-python-custom-ast.end"),
padding: 5pt,
anchor: "south",
)[P-AST#footnote[#strong[P]ython *AST*]<fn-past>]
draw.content(
"python-resolver.west",
padding: 5pt,
anchor: "south-east",
)[Resolved P-AST@fn-past]
draw.circle(
(rel: (1, -2), to: "custom-parser.south-east"),
radius: .4,
name: "midas-loader",
)
arrow(
"custom-parser",
"midas-loader",
name: "arrow-load-midas",
mark: (end: (symbol: ">", fill: black), start: "o"),
)
draw.content(
"arrow-load-midas",
anchor: "west",
padding: 5pt,
)[```python midas.using("types.midas")```]
framed(
(rel: (0, -2), to: "midas-loader.south"),
name: "midas-parser",
)[Midas lexer/parser]
arrow("midas-loader", "midas-parser", name: "arrow-midas-source")
draw.content(
"arrow-midas-source",
anchor: "west",
padding: 5pt,
)[_`types.midas`_]
framed(
(rel: (-2, 0), to: "midas-parser.west"),
anchor: "east",
name: "midas-resolver",
)[Midas Resolver]
arrow("midas-parser", "midas-resolver", name: "arrow-midas-ast")
draw.content(
"arrow-midas-ast",
anchor: "south",
padding: 5pt,
)[M-AST#footnote[#strong[M]idas *AST*]<fn-mast>]
framed(
(rel: (-3, 0), to: "midas-resolver.west"),
anchor: "east",
name: "checker",
)[Checker]
arrow("midas-resolver", "checker", name: "arrow-type-ctx")
arrow(
"python-resolver",
((), "-|", "checker.north"),
"checker",
)
draw.content(
"arrow-type-ctx",
anchor: "south",
padding: 5pt,
)[Types context]
})
#show: doc => if diagram-only {
set page(width: auto, height: auto, margin: .5cm)
diagram
} else { doc }
#align(center, title())
#v(1cm)
#figure(
diagram,
caption: [Midas type-checker architecture],
)
== Components
- *Python parser*: builtin Python AST parser, extracts abstract syntax from the raw Python source (```python ast.parse(...)```)
- *Custom python parser*: converts the raw Python AST into custom, more suitable constructs, especially for type annotations
- *Python resolver*: resolves bindings and references, tracks binding scopes
- *Midas lexer/parser*: parses a Midas type definition file and extracts its AST
- *Midas resolver*: walks the AST and fills the environment with the defined types and operations
- *Checker*: evaluates expressions and checks type coherence
@@ -0,0 +1,4 @@
a: int = 3
b: int = 4
c = a + b # -> int
@@ -0,0 +1,14 @@
type Meter(float)
type Second(float)
type MeterPerSecond(float)
extend Meter {
op __add__(Meter) -> Meter
op __sub__(Meter) -> Meter
op __truediv__(Second) -> MeterPerSecond
}
extend Second {
op __add__(Second) -> Second
op __sub__(Second) -> Second
}
@@ -0,0 +1,8 @@
# type: ignore
# ruff: disable [F821]
midas.using("02_simple_types.midas")
distance: Meter = 123.45
time: Second = 6.7
speed = distance / time
+1 -1
View File
@@ -11,7 +11,7 @@ SECTION_TEMPLATE = """{banner}
@dataclass(frozen=True, kw_only=True)
class {base}(ABC):
location: Optional[Location] = None
location: Location
@abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ...
+6 -1
View File
@@ -46,13 +46,18 @@ class Function:
args: list[Argument]
kwonlyargs: list[Argument]
returns: Optional[MidasType]
body: list[Stmt]
@dataclass(frozen=True, kw_only=True)
class Argument:
location: Optional[Location] = None
name: Optional[str]
name: str
type: Optional[MidasType]
@property
def all_args(self) -> list[Argument]:
return self.posonlyargs + self.args + self.kwonlyargs
class TypeAssign:
name: str
+2 -2
View File
@@ -21,7 +21,7 @@ T = TypeVar("T")
@dataclass(frozen=True, kw_only=True)
class Stmt(ABC):
location: Optional[Location] = None
location: Location
@abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ...
@@ -114,7 +114,7 @@ class PredicateStmt(Stmt):
@dataclass(frozen=True, kw_only=True)
class Expr(ABC):
location: Optional[Location] = None
location: Location
@abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ...
+9 -4
View File
@@ -21,7 +21,7 @@ T = TypeVar("T")
@dataclass(frozen=True, kw_only=True)
class MidasType(ABC):
location: Optional[Location] = None
location: Location
@abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ...
@@ -82,7 +82,7 @@ class FrameType(MidasType):
@dataclass(frozen=True, kw_only=True)
class Stmt(ABC):
location: Optional[Location] = None
location: Location
@abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ...
@@ -116,13 +116,18 @@ class Function(Stmt):
args: list[Argument]
kwonlyargs: list[Argument]
returns: Optional[MidasType]
body: list[Stmt]
@dataclass(frozen=True, kw_only=True)
class Argument:
location: Optional[Location] = None
name: Optional[str]
name: str
type: Optional[MidasType]
@property
def all_args(self) -> list[Argument]:
return self.posonlyargs + self.args + self.kwonlyargs
def accept(self, visitor: Stmt.Visitor[T]) -> T:
return visitor.visit_function(self)
@@ -152,7 +157,7 @@ class AssignStmt(Stmt):
@dataclass(frozen=True, kw_only=True)
class Expr(ABC):
location: Optional[Location] = None
location: Location
@abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ...
+198
View File
@@ -0,0 +1,198 @@
import logging
from pathlib import Path
from typing import Optional
import midas.ast.midas as m
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.operators import OPERATOR_METHODS
from midas.checker.types import Type, UnknownType
from midas.lexer.midas import MidasLexer
from midas.lexer.token import Token
from midas.parser.midas import MidasParser
from midas.resolver.midas import MidasResolver
class Checker(
p.Stmt.Visitor[None],
p.Expr.Visitor[Type],
p.MidasType.Visitor[Type],
):
def __init__(self, locals: dict[p.Expr, int], file_path: Path):
self.logger: logging.Logger = logging.getLogger("Checker")
self.file_path: Path = file_path
self.ctx: MidasResolver = MidasResolver()
self.global_env: Environment = Environment()
self.env: Environment = self.global_env
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:
return expr.accept(self)
def check(self, statements: list[p.Stmt]) -> list[Diagnostic]:
self.diagnostics = []
for stmt in statements:
stmt.accept(self)
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]:
distance: Optional[int] = self.locals.get(expr)
if distance is not None:
return self.env.get_at(distance, name)
return self.global_env.get(name)
def parse_midas_import(self, expr: p.CallExpr) -> Optional[Path]:
match expr:
case p.CallExpr(
callee=p.GetExpr(
object=p.VariableExpr(name="midas"),
name="using",
),
arguments=[
p.LiteralExpr(value=path),
],
):
return Path(path)
return None
def import_midas(self, path: Path) -> None:
self.logger.debug(f"Importing type definitions from {path}")
path = (self.file_path.parent / path).resolve()
lexer: MidasLexer = MidasLexer(path.read_text())
tokens: list[Token] = lexer.process()
parser: MidasParser = MidasParser(tokens)
stmts: list[m.Stmt] = parser.parse()
self.ctx.resolve(stmts)
self.logger.debug(f"Midas types: {self.ctx._types}")
self.logger.debug(f"Midas operations: {self.ctx._operations}")
def visit_expression_stmt(self, stmt: p.ExpressionStmt) -> None:
self.evaluate(stmt.expr)
def visit_function(self, stmt: p.Function) -> None: ...
def visit_type_assign(self, stmt: p.TypeAssign) -> None:
# TODO check not yet defined locally
type: Type = stmt.type.accept(self)
self.env.define(stmt.name, type)
def visit_assign_stmt(self, stmt: p.AssignStmt) -> None:
value: Type = self.evaluate(stmt.value)
for target in stmt.targets:
if not isinstance(target, p.VariableExpr):
self.logger.warning(f"Unsupported assignment to {target}")
self.warning(target.location, f"Unsupported assignment to {target}")
continue
name: str = target.name
var_type: Optional[Type] = self.look_up_variable(name, target)
if var_type is None:
self.env.define(name, value)
else:
# TODO: implement real comparison method
if var_type != value:
self.error(
stmt.location,
f"Cannot assign {value} to {name} of type {var_type}",
)
def visit_binary_expr(self, expr: p.BinaryExpr) -> Type:
method: Optional[str] = OPERATOR_METHODS.get(expr.operator.__class__)
if method is None:
self.logger.warning(f"Unsupported operator {expr.operator}")
self.warning(expr.location, f"Unsupported operator {expr.operator}")
return UnknownType()
left: Type = self.evaluate(expr.left)
right: Type = self.evaluate(expr.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
def visit_compare_expr(self, expr: p.CompareExpr) -> Type: ...
def visit_unary_expr(self, expr: p.UnaryExpr) -> Type: ...
def visit_call_expr(self, expr: p.CallExpr) -> Type:
if path := self.parse_midas_import(expr):
self.import_midas(path)
return UnknownType()
callee: Type = self.evaluate(expr.callee)
arguments: list[Type] = [self.evaluate(arg) for arg in expr.arguments]
keywords: dict[str, Type] = {
name: self.evaluate(arg) for name, arg in expr.keywords.items()
}
return UnknownType()
def visit_get_expr(self, expr: p.GetExpr) -> Type: ...
def visit_literal_expr(self, expr: p.LiteralExpr) -> Type:
match expr.value:
case bool(): # Must be before int
return self.ctx.get_type("bool")
case int():
return self.ctx.get_type("int")
case float():
return self.ctx.get_type("float")
case str():
return self.ctx.get_type("str")
case _:
self.warning(expr.location, f"Unknown literal {expr}")
return UnknownType()
def visit_variable_expr(self, expr: p.VariableExpr) -> Type:
return self.look_up_variable(expr.name, expr) or UnknownType()
def visit_logical_expr(self, expr: p.LogicalExpr) -> Type: ...
def visit_set_expr(self, expr: p.SetExpr) -> Type: ...
def visit_base_type(self, node: p.BaseType) -> Type:
return self.ctx.get_type(node.base)
def visit_constraint_type(self, node: p.ConstraintType) -> Type: ...
def visit_frame_column(self, node: p.FrameColumn) -> Type: ...
def visit_frame_type(self, node: p.FrameType) -> 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}"
+52
View File
@@ -0,0 +1,52 @@
from __future__ import annotations
from typing import Optional
from midas.checker.types import Type
class Environment:
def __init__(self, enclosing: Optional[Environment] = None) -> None:
self.enclosing: Optional[Environment] = enclosing
self.values: dict[str, Type] = {}
def define(self, name: str, value: Type):
self.values[name] = value
def get(self, name: str) -> Optional[Type]:
if name in self.values:
return self.values[name]
if self.enclosing is not None:
return self.enclosing.get(name)
# raise NameError(f"Undefined variable '{name}'")
return None
def assign(self, name: str, value: Type) -> bool:
if name not in self.values:
if self.enclosing is None:
return False
if self.enclosing.assign(name, value):
return True
self.values[name] = value
return True
def clear(self):
self.values = {}
def get_at(self, distance: int, name: str) -> Optional[Type]:
return self.ancestor(distance).values.get(name)
def assign_at(self, distance: int, name: str, value: Type):
self.ancestor(distance).values[name] = value
def ancestor(self, distance: int) -> Environment:
env: Environment = self
for _ in range(distance):
assert env.enclosing is not None
env = env.enclosing
return env
def flat_dict(self) -> dict:
if self.enclosing is None:
return self.values
return self.enclosing.flat_dict() | self.values
+31
View File
@@ -0,0 +1,31 @@
import ast
from typing import Type
OPERATOR_METHODS: dict[Type[ast.operator], str] = {
ast.Add: "__add__",
ast.Sub: "__sub__",
ast.Mult: "__mul__",
ast.MatMult: "__matmul__",
ast.Div: "__truediv__",
ast.Mod: "__mod__",
ast.Pow: "__pow__",
ast.LShift: "__lshift__",
ast.RShift: "__rshift__",
ast.BitOr: "__or__",
ast.BitXor: "__xor__",
ast.BitAnd: "__and__",
ast.FloorDiv: "__floordiv__",
}
COMPARATOR_METHODS: dict[Type[ast.cmpop], str] = {
ast.Eq: "__eq__",
# ast.NotEq: "__noteq__",
ast.Lt: "__lt__",
ast.LtE: "__le__",
ast.Gt: "__gt__",
ast.GtE: "__ge__",
# ast.Is: "__is__",
# ast.IsNot: "__isnot__",
# ast.In: "__in__",
# ast.NotIn: "__notin__",
}
+22
View File
@@ -0,0 +1,22 @@
from __future__ import annotations
from dataclasses import dataclass
@dataclass(frozen=True, kw_only=True)
class BaseType:
name: str
@dataclass(frozen=True, kw_only=True)
class SimpleType:
name: str
base: BaseType | SimpleType
@dataclass(frozen=True, kw_only=True)
class UnknownType:
pass
Type = BaseType | SimpleType | UnknownType
+20 -1
View File
@@ -1,5 +1,7 @@
import ast
import logging
from dataclasses import dataclass
from pathlib import Path
from typing import Optional, TextIO
import click
@@ -8,11 +10,14 @@ import midas.ast.midas as m
import midas.ast.python as p
from midas.ast.location import Location
from midas.ast.printer import PythonAstPrinter
from midas.checker.checker import Checker
from midas.checker.diagnostic import Diagnostic
from midas.cli.highlighter import Highlighter, MidasHighlighter, PythonHighlighter
from midas.lexer.midas import MidasLexer
from midas.lexer.token import Token, TokenType
from midas.parser.midas import MidasParser
from midas.parser.python import PythonParser
from midas.resolver.resolver import Resolver
@click.group()
@@ -23,7 +28,17 @@ def midas():
@midas.command()
@click.argument("file", type=click.File("r"))
def compile(file: TextIO):
raise NotImplementedError
logging.basicConfig(level=logging.DEBUG)
source: str = file.read()
tree: ast.Module = ast.parse(source, filename=file.name)
parser = PythonParser()
stmts: list[p.Stmt] = parser.parse_module(tree)
resolver = Resolver()
resolver.resolve(*stmts)
checker = Checker(resolver.locals, file_path=Path(file.name).resolve())
diagnostics: list[Diagnostic] = checker.check(stmts)
for diagnostic in diagnostics:
print(diagnostic)
@midas.group()
@@ -109,3 +124,7 @@ def highlight(output: TextIO, file: TextIO):
else:
raise ValueError("Unsupported file type")
highlighter.dump(output)
if __name__ == "__main__":
midas()
+6 -18
View File
@@ -205,9 +205,7 @@ class MidasParser(Parser):
while self.match(TokenType.AND):
operator: Token = self.previous()
right: Expr = self.equality()
location: Optional[Location] = None
if expr.location and right.location:
location = Location.span(expr.location, right.location)
location: Location = Location.span(expr.location, right.location)
expr = LogicalExpr(
location=location, left=expr, operator=operator, right=right
)
@@ -223,9 +221,7 @@ class MidasParser(Parser):
while self.match(TokenType.BANG_EQUAL, TokenType.EQUAL_EQUAL):
operator: Token = self.previous()
right: Expr = self.comparison()
location: Optional[Location] = None
if expr.location and right.location:
location = Location.span(expr.location, right.location)
location: Location = Location.span(expr.location, right.location)
expr = BinaryExpr(
location=location, left=expr, operator=operator, right=right
)
@@ -246,9 +242,7 @@ class MidasParser(Parser):
):
operator: Token = self.previous()
right: Expr = self.unary()
location: Optional[Location] = None
if expr.location and right.location:
location = Location.span(expr.location, right.location)
location: Location = Location.span(expr.location, right.location)
expr = BinaryExpr(
location=location, left=expr, operator=operator, right=right
)
@@ -263,9 +257,7 @@ class MidasParser(Parser):
if self.match(TokenType.MINUS):
operator: Token = self.previous()
right: Expr = self.unary()
location: Optional[Location] = None
if right.location:
location = Location.span(operator.get_location(), right.location)
location: Location = Location.span(operator.get_location(), right.location)
return UnaryExpr(location=location, operator=operator, right=right)
return self.reference()
@@ -280,9 +272,7 @@ class MidasParser(Parser):
name: Token = self.consume(
TokenType.IDENTIFIER, "Expected property name after '.'"
)
location: Optional[Location] = None
if expr.location:
location = Location.span(expr.location, name.get_location())
location: Location = Location.span(expr.location, name.get_location())
expr = GetExpr(location=location, expr=expr, name=name)
return expr
@@ -370,9 +360,7 @@ class MidasParser(Parser):
while not self.is_at_end() and not self.check(TokenType.RIGHT_BRACE):
operations.append(self.op_declaration())
self.consume(TokenType.RIGHT_BRACE, "Unclosed extend body")
location: Optional[Location] = None
if type.location:
location = keyword.location_to(self.previous())
location: Location = keyword.location_to(self.previous())
return ExtendStmt(location=location, type=type, operations=operations)
def op_declaration(self) -> OpStmt:
+43 -13
View File
@@ -53,6 +53,7 @@ class PythonParser:
return statements
def parse_stmt(self, node: ast.stmt) -> None | Stmt | list[Stmt]:
location: Location = Location.from_ast(node)
match node:
case ast.AnnAssign():
return self.parse_annotation_assign(node)
@@ -64,7 +65,10 @@ class PythonParser:
return self.parse_function(node)
case ast.Expr(value=expr):
return ExpressionStmt(expr=self.parse_expr(expr))
return ExpressionStmt(
location=location,
expr=self.parse_expr(expr),
)
case _:
print(f"Unsupported statement: {ast.unparse(node)}")
@@ -128,11 +132,19 @@ class PythonParser:
kwonlyargs=kwonlyargs,
),
returns=returns,
body=raw_body,
):
def parse_args(args_list: list[ast.arg]) -> list[Function.Argument]:
return [self._parse_function_argument(arg) for arg in args_list]
body: list[Stmt] = []
for stmt in raw_body:
stmts = self.parse_stmt(stmt)
if isinstance(stmts, Stmt):
body.append(stmts)
elif stmts is not None:
body.extend(stmts)
return Function(
location=loc,
name=name,
@@ -140,6 +152,7 @@ class PythonParser:
args=parse_args(args),
kwonlyargs=parse_args(kwonlyargs),
returns=self._parse_type(returns) if returns is not None else None,
body=body,
)
case _:
print(f"Unsupported function definition: {ast.unparse(node)}")
@@ -257,12 +270,14 @@ class PythonParser:
raise UnsupportedSyntaxError(column)
def parse_expr(self, node: ast.expr) -> Expr:
location: Location = Location.from_ast(node)
match node:
case ast.BoolOp():
return self.parse_bool_op(node)
case ast.BinOp(left=left, op=op, right=right):
return BinaryExpr(
location=location,
left=self.parse_expr(left),
operator=op,
right=self.parse_expr(right),
@@ -270,6 +285,7 @@ class PythonParser:
case ast.UnaryOp(op=op, operand=right):
return UnaryExpr(
location=location,
operator=op,
right=self.parse_expr(right),
)
@@ -281,33 +297,39 @@ class PythonParser:
return self.parse_call(node)
case ast.Constant(value=value):
return LiteralExpr(value=value)
return LiteralExpr(location=location, value=value)
case ast.Attribute(value=object, attr=name):
return GetExpr(
location=location,
object=self.parse_expr(object),
name=name,
)
case ast.Name(id=name):
return VariableExpr(name=name)
return VariableExpr(location=location, 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
rights: list[Expr] = [self.parse_expr(expr) for expr in node.values]
expr: LogicalExpr = LogicalExpr(
left=self.parse_expr(values[0]),
location=Location.span(
rights[0].location,
rights[1].location,
),
left=rights[0],
operator=op,
right=self.parse_expr(values[1]),
right=rights[1],
)
for value in values[2:]:
for right in rights[2:]:
expr = LogicalExpr(
location=Location.span(expr.location, right.location),
left=expr,
operator=op,
right=self.parse_expr(value),
right=right,
)
return expr
@@ -315,24 +337,32 @@ class PythonParser:
ops: list[ast.cmpop] = node.ops
rights: list[Expr] = [self.parse_expr(expr) for expr in node.comparators]
expr: Expr = CompareExpr(
location=Location.span(
rights[0].location,
rights[1].location,
),
left=self.parse_expr(node.left),
operator=ops[0],
right=rights[0],
)
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(
location=Location.span(expr.location, comparison.location),
left=expr,
operator=ast.And(),
right=CompareExpr(
left=rights[i],
operator=ops[i],
right=right,
),
right=comparison,
)
return expr
def parse_call(self, node: ast.Call) -> CallExpr:
return CallExpr(
location=Location.from_ast(node),
callee=self.parse_expr(node.func),
arguments=[self.parse_expr(arg) for arg in node.args],
keywords={
+113
View File
@@ -0,0 +1,113 @@
from typing import Optional
import midas.ast.midas as m
from midas.checker.types import BaseType, SimpleType, Type
class MidasResolver(m.Stmt.Visitor[None], m.Expr.Visitor[Type]):
def __init__(self) -> None:
self._types: dict[str, Type] = {}
self._operations: dict[tuple[Type, str, Type], Type] = {}
self._define_builtin()
def get_type(self, name: str) -> Type:
type: Optional[Type] = self._types.get(name)
if type is None:
raise NameError(f"Undefined type {name}")
return type
def get_operation_result(
self, left: Type, operator: str, right: Type
) -> Optional[Type]:
operation: tuple[Type, str, Type] = (left, operator, right)
result: Optional[Type] = self._operations.get(operation)
return result
def _define_builtin(self):
self.define_type("bool", BaseType(name="bool"))
self.define_type("int", BaseType(name="int"))
self.define_type("float", BaseType(name="float"))
self.define_type("str", BaseType(name="str"))
self.define_operation(
left=self.get_type("int"),
operator="__add__",
right=self.get_type("int"),
result=self.get_type("int"),
)
def define_type(self, name: str, type: Type) -> Type:
if name in self._types:
raise ValueError(f"Type {name} already defined")
self._types[name] = type
return type
def define_operation(self, left: Type, operator: str, right: Type, result: Type):
operation: tuple[Type, str, Type] = (left, operator, right)
if operation in self._operations:
raise ValueError(
f"Operation {operator} already defined between {left} and {right}"
)
self._operations[operation] = result
def resolve(self, stmts: list[m.Stmt]):
for stmt in stmts:
stmt.accept(self)
def visit_simple_type_stmt(self, stmt: m.SimpleTypeStmt) -> None:
# TODO generics, optional, constraint
base: Type = self.get_type(stmt.base.name.lexeme)
match base:
case BaseType() | SimpleType():
type = SimpleType(
name=stmt.name.lexeme,
base=base,
)
self.define_type(type.name, type)
case _:
raise TypeError(f"Invalid base {base} for simple type")
def visit_complex_type_stmt(self, stmt: m.ComplexTypeStmt) -> None: ...
def visit_property_stmt(self, stmt: m.PropertyStmt) -> None: ...
def visit_extend_stmt(self, stmt: m.ExtendStmt) -> None:
base: Type = stmt.type.accept(self)
for op in stmt.operations:
right: Type = op.operand.accept(self)
result: Type = op.result.accept(self)
self.define_operation(
left=base,
operator=op.name.lexeme,
right=right,
result=result,
)
def visit_op_stmt(self, stmt: m.OpStmt) -> None: ...
def visit_predicate_stmt(self, stmt: m.PredicateStmt) -> None: ...
def visit_simple_type_expr(self, expr: m.SimpleTypeExpr) -> Type:
return self.get_type(expr.name.lexeme)
def visit_logical_expr(self, expr: m.LogicalExpr) -> Type: ...
def visit_binary_expr(self, expr: m.BinaryExpr) -> Type: ...
def visit_unary_expr(self, expr: m.UnaryExpr) -> Type: ...
def visit_get_expr(self, expr: m.GetExpr) -> Type: ...
def visit_variable_expr(self, expr: m.VariableExpr) -> Type: ...
def visit_grouping_expr(self, expr: m.GroupingExpr) -> Type:
return expr.expr.accept(self)
def visit_literal_expr(self, expr: m.LiteralExpr) -> Type: ...
def visit_wildcard_expr(self, expr: m.WildcardExpr) -> Type: ...
def visit_template_expr(self, expr: m.TemplateExpr) -> Type: ...
def visit_type_expr(self, expr: m.TypeExpr) -> Type:
return self.get_type(expr.name.lexeme)
+112
View File
@@ -0,0 +1,112 @@
import midas.ast.python as p
class ResolverError(Exception): ...
class Resolver(p.Stmt.Visitor[None], p.Expr.Visitor[None]):
def __init__(self):
self.locals: dict[p.Expr, int] = {}
self.scopes: list[dict[str, bool]] = []
def resolve(self, *objects: p.Stmt | p.Expr) -> None:
for obj in objects:
obj.accept(self)
def begin_scope(self):
self.scopes.append({})
def end_scope(self):
self.scopes.pop()
def declare(self, name: str) -> None:
if len(self.scopes) == 0:
return
scope: dict[str, bool] = self.scopes[-1]
if name in scope:
raise ResolverError(
f"A variable with the name {name} is already declared in this scope"
)
scope[name] = False
def define(self, name: str) -> None:
if len(self.scopes) == 0:
return
self.scopes[-1][name] = True
def resolve_local(self, expr: p.Expr, name: str) -> None:
for i, scope in enumerate(reversed(self.scopes)):
if name in scope:
self.locals[expr] = i
return
def resolve_function(self, function: p.Function) -> None:
self.begin_scope()
for param in function.all_args:
self.declare(param.name)
self.define(param.name)
self.resolve(*function.body)
self.end_scope()
def visit_expression_stmt(self, stmt: p.ExpressionStmt) -> None:
stmt.expr.accept(self)
def visit_function(self, stmt: p.Function) -> None:
# Declare before resolving body to allow recursion
self.declare(stmt.name)
self.define(stmt.name)
self.resolve_function(stmt)
def visit_type_assign(self, stmt: p.TypeAssign) -> None:
self.declare(stmt.name)
# NOTE: resolve type here?
self.define(stmt.name)
def visit_assign_stmt(self, stmt: p.AssignStmt) -> None:
self.resolve(stmt.value)
for target in stmt.targets:
match target:
case p.VariableExpr(name=name):
self.resolve_local(target, name)
# TODO: declare if not found
case _:
raise Exception(f"Unsupported assignment to {target}")
def visit_binary_expr(self, expr: p.BinaryExpr) -> None:
self.resolve(expr.left)
self.resolve(expr.right)
def visit_compare_expr(self, expr: p.CompareExpr) -> None:
self.resolve(expr.left)
self.resolve(expr.right)
def visit_unary_expr(self, expr: p.UnaryExpr) -> None:
self.resolve(expr.right)
def visit_call_expr(self, expr: p.CallExpr) -> None:
self.resolve(expr.callee)
for arg in expr.arguments:
self.resolve(arg)
for arg in expr.keywords.values():
self.resolve(arg)
def visit_get_expr(self, expr: p.GetExpr) -> None:
self.resolve(expr.object)
def visit_literal_expr(self, expr: p.LiteralExpr) -> None:
pass
def visit_variable_expr(self, expr: p.VariableExpr) -> None:
if len(self.scopes) != 0 and self.scopes[-1].get(expr.name) is False:
raise ResolverError(
f"Cannot use local variable '{expr.name}' in its own initializer"
) # aka. UnboundLocalError
self.resolve_local(expr, expr.name)
def visit_logical_expr(self, expr: p.LogicalExpr) -> None:
self.resolve(expr.left)
self.resolve(expr.right)
def visit_set_expr(self, expr: p.SetExpr) -> None:
self.resolve(expr.value)
self.resolve(expr.object)