Compare commits

...
3 Commits
Author SHA1 Message Date
HEL a735113466 fix(parser): update ast gen script 2026-05-25 12:46:04 +02:00
HEL 0e0a1b26f2 feat(cli): add midas highlighter 2026-05-25 12:14:55 +02:00
HEL e94db2181f feat(parser): add location to midas AST nodes 2026-05-25 12:14:14 +02:00
12 changed files with 439 additions and 118 deletions

No files matched your search

+8 -3
View File
@@ -14,7 +14,8 @@ from abc import ABC, abstractmethod
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Generic, Optional, TypeVar from typing import Any, Generic, Optional, TypeVar
from lexer.token import Token from midas.ast.location import Location
from midas.lexer.token import Token
T = TypeVar("T") T = TypeVar("T")
@@ -23,8 +24,10 @@ T = TypeVar("T")
############## ##############
@dataclass(frozen=True) @dataclass(frozen=True, kw_only=True)
class Stmt(ABC): class Stmt(ABC):
location: Optional[Location] = None
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
@@ -40,8 +43,10 @@ class Stmt(ABC):
############### ###############
@dataclass(frozen=True) @dataclass(frozen=True, kw_only=True)
class Expr(ABC): class Expr(ABC):
location: Optional[Location] = None
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
+37
View File
@@ -0,0 +1,37 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Optional, Protocol
class HasLocation(Protocol):
lineno: int
col_offset: int
end_lineno: Optional[int]
end_col_offset: Optional[int]
@dataclass(frozen=True, kw_only=True)
class Location:
lineno: int
col_offset: int
end_lineno: Optional[int]
end_col_offset: Optional[int]
@staticmethod
def from_ast(obj: HasLocation) -> Location:
return Location(
lineno=obj.lineno,
col_offset=obj.col_offset,
end_lineno=obj.end_lineno,
end_col_offset=obj.end_col_offset,
)
@staticmethod
def span(start: Location, end: Location) -> Location:
return Location(
lineno=start.lineno,
col_offset=start.col_offset,
end_lineno=end.lineno,
end_col_offset=end.end_col_offset,
)
+7 -2
View File
@@ -9,6 +9,7 @@ from abc import ABC, abstractmethod
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Generic, Optional, TypeVar from typing import Any, Generic, Optional, TypeVar
from midas.ast.location import Location
from midas.lexer.token import Token from midas.lexer.token import Token
T = TypeVar("T") T = TypeVar("T")
@@ -18,8 +19,10 @@ T = TypeVar("T")
############## ##############
@dataclass(frozen=True) @dataclass(frozen=True, kw_only=True)
class Stmt(ABC): class Stmt(ABC):
location: Optional[Location] = None
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
@@ -109,8 +112,10 @@ class PredicateStmt(Stmt):
############### ###############
@dataclass(frozen=True) @dataclass(frozen=True, kw_only=True)
class Expr(ABC): class Expr(ABC):
location: Optional[Location] = None
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
+3 -25
View File
@@ -3,35 +3,13 @@ from __future__ import annotations
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
import ast import ast
from dataclasses import dataclass from dataclasses import dataclass
from typing import Generic, Optional, Protocol, TypeVar from typing import Generic, Optional, TypeVar
from midas.ast.location import Location
T = TypeVar("T") T = TypeVar("T")
class HasLocation(Protocol):
lineno: int
col_offset: int
end_lineno: Optional[int]
end_col_offset: Optional[int]
@dataclass(frozen=True, kw_only=True)
class Location:
lineno: int
col_offset: int
end_lineno: Optional[int]
end_col_offset: Optional[int]
@staticmethod
def from_ast(obj: HasLocation) -> Location:
return Location(
lineno=obj.lineno,
col_offset=obj.col_offset,
end_lineno=obj.end_lineno,
end_col_offset=obj.end_col_offset,
)
@dataclass(frozen=True, kw_only=True) @dataclass(frozen=True, kw_only=True)
class Expr(ABC): class Expr(ABC):
location: Optional[Location] = None location: Optional[Location] = None
-28
View File
@@ -50,32 +50,4 @@ span {
--border: 2px; --border: 2px;
z-index: 10; z-index: 10;
} }
&.base-type {
--col: 108, 233, 108;
}
&.param {
--col: 103, 192, 224;
}
&.constraint-type {
--col: 174, 200, 195;
}
&.frame-column {
--col: 216, 231, 81;
}
&.frame-type {
--col: 231, 46, 40;
}
&.function {
--col: 215, 103, 224;
}
&.argument {
--col: 103, 192, 224;
}
} }
+138 -26
View File
@@ -1,19 +1,29 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from pathlib import Path from pathlib import Path
from typing import TextIO from typing import Generic, Optional, Protocol, TextIO, TypeVar
from midas.ast.python import ( from midas.ast.location import Location
BaseType, import midas.ast.midas as m
ConstraintType, import midas.ast.python as p
Expr,
FrameColumn, H = TypeVar("H", bound="Highlighter", contravariant=True)
FrameType,
Function,
FunctionArgument,
)
class PythonHighlighter(Expr.Visitor[None]): class Highlightable(Protocol, Generic[H]):
CSS_PATH: Path = Path(__file__).parent / "highlight.css" def accept(self, visitor: H): ...
class Locatable(Protocol):
@property
@abstractmethod
def location(self) -> Optional[Location]: ...
class Highlighter(ABC):
BASE_CSS_PATH: Path = Path(__file__).parent / "highlight.css"
EXTRA_CSS_PATH: Optional[Path] = None
def __init__(self, source: str) -> None: def __init__(self, source: str) -> None:
self.source: str = source self.source: str = source
@@ -21,12 +31,22 @@ class PythonHighlighter(Expr.Visitor[None]):
self.openings: dict[tuple[int, int], list[str]] = {} self.openings: dict[tuple[int, int], list[str]] = {}
self.closings: dict[tuple[int, int], list[str]] = {} self.closings: dict[tuple[int, int], list[str]] = {}
def highlight(self, node: Expr): def format_css(self, path: Path) -> list[str]:
node.accept(self) css: str = path.read_text()
css = "\n".join((" " + line).rstrip() for line in css.splitlines())
return [
" <style>",
css,
" </style>",
]
def dump(self, buf: TextIO): def dump(self, buf: TextIO):
css: str = self.CSS_PATH.read_text() base_css: list[str] = self.format_css(self.BASE_CSS_PATH)
css = "\n".join((" " + line).rstrip() for line in css.splitlines()) extra_css: list[str] = (
self.format_css(self.EXTRA_CSS_PATH)
if self.EXTRA_CSS_PATH is not None
else []
)
lines: list[str] = [ lines: list[str] = [
"<!DOCTYPE html>", "<!DOCTYPE html>",
'<html lang="en">', '<html lang="en">',
@@ -34,9 +54,8 @@ class PythonHighlighter(Expr.Visitor[None]):
' <meta charset="UTF-8">', ' <meta charset="UTF-8">',
' <meta name="viewport" content="width=device-width, initial-scale=1.0">', ' <meta name="viewport" content="width=device-width, initial-scale=1.0">',
" <title>Highlighted file</title>", " <title>Highlighted file</title>",
" <style>", *base_css,
css, *extra_css,
" </style>",
"</head>", "</head>",
"<body>", "<body>",
' <div id="code">', ' <div id="code">',
@@ -64,7 +83,7 @@ class PythonHighlighter(Expr.Visitor[None]):
buf.write("\n".join(lines)) buf.write("\n".join(lines))
def wrap(self, node: Expr, cls: str): def wrap(self, node: Locatable, cls: str):
if node.location is None: if node.location is None:
return return
if node.location.end_lineno is None or node.location.end_col_offset is None: if node.location.end_lineno is None or node.location.end_col_offset is None:
@@ -84,32 +103,125 @@ class PythonHighlighter(Expr.Visitor[None]):
self.closings.setdefault((l, c), []).insert(0, closing) self.closings.setdefault((l, c), []).insert(0, closing)
self.openings.setdefault((l + 1, 0), []).append(opening) self.openings.setdefault((l + 1, 0), []).append(opening)
def visit_base_type(self, node: BaseType) -> None:
class PythonHighlighter(Highlighter, p.Expr.Visitor[None]):
EXTRA_CSS_PATH: Optional[Path] = Path(__file__).parent / "hl_python.css"
def highlight(self, node: Highlightable[PythonHighlighter]):
node.accept(self)
def visit_base_type(self, node: p.BaseType) -> None:
self.wrap(node, "base-type") self.wrap(node, "base-type")
if node.param is not None: if node.param is not None:
self.wrap(node.param, "param") self.wrap(node.param, "param")
node.param.accept(self) node.param.accept(self)
def visit_constraint_type(self, node: ConstraintType) -> None: def visit_constraint_type(self, node: p.ConstraintType) -> None:
self.wrap(node, "constraint-type") self.wrap(node, "constraint-type")
node.type.accept(self) node.type.accept(self)
def visit_frame_column(self, node: FrameColumn) -> None: def visit_frame_column(self, node: p.FrameColumn) -> None:
self.wrap(node, "frame-column") self.wrap(node, "frame-column")
if node.type is not None: if node.type is not None:
node.type.accept(self) node.type.accept(self)
def visit_frame_type(self, node: FrameType) -> None: def visit_frame_type(self, node: p.FrameType) -> None:
self.wrap(node, "frame-type") self.wrap(node, "frame-type")
for column in node.columns: for column in node.columns:
column.accept(self) column.accept(self)
def visit_function(self, node: Function) -> None: def visit_function(self, node: p.Function) -> None:
self.wrap(node, "function") self.wrap(node, "function")
for arg in node.posonlyargs + node.args + node.kwonlyargs: for arg in node.posonlyargs + node.args + node.kwonlyargs:
arg.accept(self) arg.accept(self)
def visit_function_argument(self, node: FunctionArgument) -> None: def visit_function_argument(self, node: p.FunctionArgument) -> None:
self.wrap(node, "argument") self.wrap(node, "argument")
if node.type is not None: if node.type is not None:
node.type.accept(self) node.type.accept(self)
class MidasHighlighter(Highlighter, m.Stmt.Visitor[None], m.Expr.Visitor[None]):
EXTRA_CSS_PATH: Optional[Path] = Path(__file__).parent / "hl_midas.css"
def highlight(self, node: Highlightable[MidasHighlighter]):
node.accept(self)
def visit_simple_type_stmt(self, stmt: m.SimpleTypeStmt) -> None:
self.wrap(stmt, "simple-type")
if stmt.template is not None:
stmt.template.accept(self)
stmt.base.accept(self)
if stmt.constraint is not None:
self.wrap(stmt.constraint, "constraint")
stmt.constraint.accept(self)
def visit_complex_type_stmt(self, stmt: m.ComplexTypeStmt) -> None:
self.wrap(stmt, "complex-type")
if stmt.template is not None:
stmt.template.accept(self)
for prop in stmt.properties:
prop.accept(self)
def visit_property_stmt(self, stmt: m.PropertyStmt) -> None:
self.wrap(stmt, "property")
stmt.type.accept(self)
if stmt.constraint is not None:
self.wrap(stmt.constraint, "constraint")
stmt.constraint.accept(self)
def visit_extend_stmt(self, stmt: m.ExtendStmt) -> None:
self.wrap(stmt, "extend")
stmt.type.accept(self)
for op in stmt.operations:
op.accept(self)
def visit_op_stmt(self, stmt: m.OpStmt) -> None:
self.wrap(stmt, "op")
stmt.operand.accept(self)
stmt.result.accept(self)
def visit_predicate_stmt(self, stmt: m.PredicateStmt) -> None:
self.wrap(stmt, "predicate")
stmt.type.accept(self)
stmt.condition.accept(self)
def visit_simple_type_expr(self, expr: m.SimpleTypeExpr) -> None:
self.wrap(expr, "simple-type-expr")
def visit_logical_expr(self, expr: m.LogicalExpr) -> None:
self.wrap(expr, "logical-expr")
expr.left.accept(self)
expr.right.accept(self)
def visit_binary_expr(self, expr: m.BinaryExpr) -> None:
self.wrap(expr, "binary-expr")
expr.left.accept(self)
expr.right.accept(self)
def visit_unary_expr(self, expr: m.UnaryExpr) -> None:
self.wrap(expr, "unary-expr")
expr.right.accept(self)
def visit_get_expr(self, expr: m.GetExpr) -> None:
self.wrap(expr, "get-expr")
expr.expr.accept(self)
def visit_variable_expr(self, expr: m.VariableExpr) -> None:
self.wrap(expr, "variable")
def visit_grouping_expr(self, expr: m.GroupingExpr) -> None:
expr.expr.accept(self)
def visit_literal_expr(self, expr: m.LiteralExpr) -> None: ...
def visit_wildcard_expr(self, expr: m.WildcardExpr) -> None: ...
def visit_template_expr(self, expr: m.TemplateExpr) -> None:
self.wrap(expr, "template")
expr.type.accept(self)
def visit_type_expr(self, expr: m.TypeExpr) -> None:
self.wrap(expr, "type")
if expr.template is not None:
expr.template.accept(self)
+55
View File
@@ -0,0 +1,55 @@
span {
&.comment {
--col: 200, 200, 200;
color: rgb(110, 110, 110);
font-style: italic;
}
&.simple-type {
--col: 108, 233, 108;
}
&.complex-type {
--col: 233, 206, 108;
}
&.constraint {
--col: 233, 108, 108;
}
&.property {
--col: 233, 108, 176;
}
&.extend {
--col: 108, 197, 233;
}
&.op {
--col: 108, 148, 233;
}
&.predicate {
--col: 193, 108, 233;
}
&.simple-type-expr {
--col: 150, 150, 150;
}
&.logical-expr,
&.binary-expr,
&.unary-expr,
&.get-expr {
--col: 123, 215, 193;
}
&.template {
--col: 163, 117, 71;
}
&.type {
--col: 200, 200, 200;
font-weight: bold;
}
}
+29
View File
@@ -0,0 +1,29 @@
span {
&.base-type {
--col: 108, 233, 108;
}
&.param {
--col: 103, 192, 224;
}
&.constraint-type {
--col: 174, 200, 195;
}
&.frame-column {
--col: 216, 231, 81;
}
&.frame-type {
--col: 231, 46, 40;
}
&.function {
--col: 215, 103, 224;
}
&.argument {
--col: 103, 192, 224;
}
}
+48 -8
View File
@@ -3,8 +3,13 @@ from typing import Optional, TextIO
import click import click
from midas.ast.location import Location
import midas.ast.midas as m
from midas.ast.printer import PythonAstPrinter from midas.ast.printer import PythonAstPrinter
from midas.cli.highlighter import PythonHighlighter 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.parser.python import PythonParser
@@ -59,18 +64,53 @@ def dump_ast(output: Optional[TextIO], parse: bool, file: TextIO):
output.write(dump) output.write(dump)
@utils.command() def highlight_python(source: str, path: str) -> Highlighter:
@click.option("-o", "--output", type=click.File("w"), default="-") tree: ast.Module = ast.parse(source, filename=path)
@click.argument("file", type=click.File("r"))
def highlight(output: TextIO, file: TextIO):
source: str = file.read()
tree: ast.Module = ast.parse(source, filename=file.name)
parser = PythonParser() parser = PythonParser()
parser.visit(tree) parser.visit(tree)
highlighter: PythonHighlighter = PythonHighlighter(source) highlighter = PythonHighlighter(source)
for _, annotation in parser.annotations: for _, annotation in parser.annotations:
if annotation is not None: if annotation is not None:
highlighter.highlight(annotation) highlighter.highlight(annotation)
for func in parser.functions: for func in parser.functions:
highlighter.highlight(func) highlighter.highlight(func)
return highlighter
def highlight_midas(source: str, path: str) -> Highlighter:
lexer = MidasLexer(source, file=path)
tokens: list[Token] = lexer.process()
parser = MidasParser(tokens)
stmts: list[m.Stmt] = parser.parse()
highlighter = MidasHighlighter(source)
class LocatableToken:
def __init__(self, token: Token):
self.token: Token = token
@property
def location(self) -> Location:
return self.token.get_location()
for token in tokens:
if token.type == TokenType.COMMENT:
highlighter.wrap(LocatableToken(token), "comment")
for stmt in stmts:
highlighter.highlight(stmt)
return highlighter
@utils.command()
@click.option("-o", "--output", type=click.File("w"), default="-")
@click.argument("file", type=click.File("r"))
def highlight(output: TextIO, file: TextIO):
source: str = file.read()
highlighter: Highlighter
if file.name.endswith(".py"):
highlighter = highlight_python(source, file.name)
elif file.name.endswith(".midas"):
highlighter = highlight_midas(source, file.name)
else:
raise ValueError("Unsupported file type")
highlighter.dump(output) highlighter.dump(output)
+23
View File
@@ -1,7 +1,10 @@
from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from enum import Enum, auto from enum import Enum, auto
from typing import Any from typing import Any
from midas.ast.location import Location
from midas.lexer.position import Position from midas.lexer.position import Position
@@ -63,3 +66,23 @@ class Token:
lexeme: str lexeme: str
value: Any value: Any
position: Position position: Position
def get_location(self) -> Location:
lineno: int = self.position.line
col_offset: int = self.position.column - 1
end_lineno = lineno
end_col_offset = col_offset
for c in self.lexeme:
end_col_offset += 1
if c == "\n":
end_lineno += 1
end_col_offset = 0
return Location(
lineno=lineno,
col_offset=col_offset,
end_lineno=end_lineno,
end_col_offset=end_col_offset,
)
def location_to(self, to: Token) -> Location:
return Location.span(self.get_location(), to.get_location())
+90 -25
View File
@@ -1,5 +1,6 @@
from typing import Optional from typing import Optional
from midas.ast.location import Location
from midas.ast.midas import ( from midas.ast.midas import (
BinaryExpr, BinaryExpr,
ComplexTypeStmt, ComplexTypeStmt,
@@ -104,6 +105,7 @@ class MidasParser(Parser):
Returns: Returns:
TypeStmt: the parsed type declaration statement TypeStmt: the parsed type declaration statement
""" """
keyword: Token = self.previous()
name: Token = self.consume(TokenType.IDENTIFIER, "Expected type name") name: Token = self.consume(TokenType.IDENTIFIER, "Expected type name")
template: Optional[TemplateExpr] = None template: Optional[TemplateExpr] = None
if self.check(TokenType.LEFT_BRACKET): if self.check(TokenType.LEFT_BRACKET):
@@ -116,11 +118,20 @@ class MidasParser(Parser):
if self.match(TokenType.WHERE): if self.match(TokenType.WHERE):
constraint = self.constraint() constraint = self.constraint()
return SimpleTypeStmt( return SimpleTypeStmt(
name=name, template=template, base=base, constraint=constraint location=keyword.location_to(self.previous()),
name=name,
template=template,
base=base,
constraint=constraint,
) )
else: else:
properties: list[PropertyStmt] = self.type_properties() properties: list[PropertyStmt] = self.type_properties()
return ComplexTypeStmt(name=name, template=template, properties=properties) return ComplexTypeStmt(
location=keyword.location_to(self.previous()),
name=name,
template=template,
properties=properties,
)
def template_expr(self) -> TemplateExpr: def template_expr(self) -> TemplateExpr:
"""Parse a generic template expression """Parse a generic template expression
@@ -130,10 +141,14 @@ class MidasParser(Parser):
Returns: Returns:
TemplateExpr: the parsed template expression TemplateExpr: the parsed template expression
""" """
self.consume(TokenType.LEFT_BRACKET, "Missing '[' before template expression") left: Token = self.consume(
TokenType.LEFT_BRACKET, "Missing '[' before template expression"
)
type: TypeExpr = self.type_expr() type: TypeExpr = self.type_expr()
self.consume(TokenType.RIGHT_BRACKET, "Missing ']' after template expression") right: Token = self.consume(
return TemplateExpr(type=type) TokenType.RIGHT_BRACKET, "Missing ']' after template expression"
)
return TemplateExpr(location=left.location_to(right), type=type)
def type_expr(self) -> TypeExpr: def type_expr(self) -> TypeExpr:
"""Parse a type expression """Parse a type expression
@@ -149,7 +164,12 @@ class MidasParser(Parser):
if self.check(TokenType.LEFT_BRACKET): if self.check(TokenType.LEFT_BRACKET):
template = self.template_expr() template = self.template_expr()
optional: bool = self.match(TokenType.QMARK) optional: bool = self.match(TokenType.QMARK)
return TypeExpr(name=name, template=template, optional=optional) return TypeExpr(
location=name.location_to(self.previous()),
name=name,
template=template,
optional=optional,
)
def simple_type_expr(self) -> SimpleTypeExpr: def simple_type_expr(self) -> SimpleTypeExpr:
"""Parse a simple type expression """Parse a simple type expression
@@ -161,7 +181,9 @@ class MidasParser(Parser):
""" """
name: Token = self.consume(TokenType.IDENTIFIER, "Expected type name") name: Token = self.consume(TokenType.IDENTIFIER, "Expected type name")
optional: bool = self.match(TokenType.QMARK) optional: bool = self.match(TokenType.QMARK)
return SimpleTypeExpr(name=name, optional=optional) return SimpleTypeExpr(
location=name.location_to(self.previous()), name=name, optional=optional
)
def constraint(self) -> Expr: def constraint(self) -> Expr:
"""Parse a constraint """Parse a constraint
@@ -183,7 +205,12 @@ 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()
expr = LogicalExpr(left=expr, operator=operator, right=right) location: Optional[Location] = None
if expr.location and right.location:
location = Location.span(expr.location, right.location)
expr = LogicalExpr(
location=location, left=expr, operator=operator, right=right
)
return expr return expr
def equality(self) -> Expr: def equality(self) -> Expr:
@@ -196,7 +223,12 @@ 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()
expr = BinaryExpr(left=expr, operator=operator, right=right) location: Optional[Location] = None
if expr.location and right.location:
location = Location.span(expr.location, right.location)
expr = BinaryExpr(
location=location, left=expr, operator=operator, right=right
)
return expr return expr
def comparison(self) -> Expr: def comparison(self) -> Expr:
@@ -214,7 +246,12 @@ class MidasParser(Parser):
): ):
operator: Token = self.previous() operator: Token = self.previous()
right: Expr = self.unary() right: Expr = self.unary()
expr = BinaryExpr(left=expr, operator=operator, right=right) location: Optional[Location] = None
if expr.location and right.location:
location = Location.span(expr.location, right.location)
expr = BinaryExpr(
location=location, left=expr, operator=operator, right=right
)
return expr return expr
def unary(self) -> Expr: def unary(self) -> Expr:
@@ -226,7 +263,10 @@ 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()
return UnaryExpr(operator=operator, right=right) location: Optional[Location] = None
if right.location:
location = Location.span(operator.get_location(), right.location)
return UnaryExpr(location=location, operator=operator, right=right)
return self.reference() return self.reference()
def reference(self) -> Expr: def reference(self) -> Expr:
@@ -240,7 +280,10 @@ class MidasParser(Parser):
name: Token = self.consume( name: Token = self.consume(
TokenType.IDENTIFIER, "Expected property name after '.'" TokenType.IDENTIFIER, "Expected property name after '.'"
) )
expr = GetExpr(expr=expr, name=name) location: Optional[Location] = None
if expr.location:
location = Location.span(expr.location, name.get_location())
expr = GetExpr(location=location, expr=expr, name=name)
return expr return expr
def primary(self) -> Expr: def primary(self) -> Expr:
@@ -251,26 +294,27 @@ class MidasParser(Parser):
Returns: Returns:
Expr: the parsed expression Expr: the parsed expression
""" """
token: Token = self.peek()
if self.match(TokenType.FALSE): if self.match(TokenType.FALSE):
return LiteralExpr(False) return LiteralExpr(location=token.get_location(), value=False)
if self.match(TokenType.TRUE): if self.match(TokenType.TRUE):
return LiteralExpr(True) return LiteralExpr(location=token.get_location(), value=True)
if self.match(TokenType.NONE): if self.match(TokenType.NONE):
return LiteralExpr(None) return LiteralExpr(location=token.get_location(), value=None)
if self.match(TokenType.NUMBER): if self.match(TokenType.NUMBER):
return LiteralExpr(self.previous().value) return LiteralExpr(location=token.get_location(), value=token.value)
if self.match(TokenType.IDENTIFIER): if self.match(TokenType.IDENTIFIER):
return VariableExpr(self.previous()) return VariableExpr(location=token.get_location(), name=token)
if self.match(TokenType.UNDERSCORE): if self.match(TokenType.UNDERSCORE):
return WildcardExpr(self.previous()) return WildcardExpr(location=token.get_location(), token=token)
if self.match(TokenType.LEFT_PAREN): if self.match(TokenType.LEFT_PAREN):
expr: Expr = self.constraint() expr: Expr = self.constraint()
self.consume(TokenType.RIGHT_PAREN, "Unclosed parenthesis") right: Token = self.consume(TokenType.RIGHT_PAREN, "Unclosed parenthesis")
return GroupingExpr(expr) return GroupingExpr(location=token.location_to(right), expr=expr)
raise self.error(self.peek(), "Expected expression") raise self.error(self.peek(), "Expected expression")
@@ -304,7 +348,12 @@ class MidasParser(Parser):
constraint: Optional[Expr] = None constraint: Optional[Expr] = None
if self.match(TokenType.WHERE): if self.match(TokenType.WHERE):
constraint = self.constraint() constraint = self.constraint()
return PropertyStmt(name=name, type=type, constraint=constraint) return PropertyStmt(
location=name.location_to(self.previous()),
name=name,
type=type,
constraint=constraint,
)
def extend_declaration(self) -> ExtendStmt: def extend_declaration(self) -> ExtendStmt:
"""Parse an extension definition """Parse an extension definition
@@ -314,13 +363,17 @@ class MidasParser(Parser):
Returns: Returns:
ExtendStmt: the parsed extension statement ExtendStmt: the parsed extension statement
""" """
keyword: Token = self.previous()
type: TypeExpr = self.type_expr() type: TypeExpr = self.type_expr()
self.consume(TokenType.LEFT_BRACE, "Expected '{' to start extend body") self.consume(TokenType.LEFT_BRACE, "Expected '{' to start extend body")
operations: list[OpStmt] = [] operations: list[OpStmt] = []
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")
return ExtendStmt(type=type, operations=operations) location: Optional[Location] = None
if type.location:
location = keyword.location_to(self.previous())
return ExtendStmt(location=location, type=type, operations=operations)
def op_declaration(self) -> OpStmt: def op_declaration(self) -> OpStmt:
"""Parse an operation definition """Parse an operation definition
@@ -330,7 +383,7 @@ class MidasParser(Parser):
Returns: Returns:
OpStmt: the parsed operation statement OpStmt: the parsed operation statement
""" """
self.consume(TokenType.OP, "Expected 'op' keyword") keyword: Token = self.consume(TokenType.OP, "Expected 'op' keyword")
name: Token = self.consume(TokenType.IDENTIFIER, "Expected operation name") name: Token = self.consume(TokenType.IDENTIFIER, "Expected operation name")
self.consume(TokenType.LEFT_PAREN, "Expected '(' before operand type") self.consume(TokenType.LEFT_PAREN, "Expected '(' before operand type")
@@ -340,7 +393,12 @@ class MidasParser(Parser):
self.consume(TokenType.ARROW, "Expected '->' before result type") self.consume(TokenType.ARROW, "Expected '->' before result type")
result: TypeExpr = self.type_expr() result: TypeExpr = self.type_expr()
return OpStmt(name=name, operand=operand, result=result) return OpStmt(
location=keyword.location_to(self.previous()),
name=name,
operand=operand,
result=result,
)
def predicate_declaration(self) -> PredicateStmt: def predicate_declaration(self) -> PredicateStmt:
"""Parse a predicate declaration """Parse a predicate declaration
@@ -350,6 +408,7 @@ class MidasParser(Parser):
Returns: Returns:
PredicateStmt: the parsed predicate declaration statement PredicateStmt: the parsed predicate declaration statement
""" """
keyword: Token = self.previous()
name: Token = self.consume(TokenType.IDENTIFIER, "Expected predicate name") name: Token = self.consume(TokenType.IDENTIFIER, "Expected predicate name")
self.consume(TokenType.LEFT_PAREN, "Expected '(' before predicate subject") self.consume(TokenType.LEFT_PAREN, "Expected '(' before predicate subject")
subject: Token = self.consume(TokenType.IDENTIFIER, "Expected subject name") subject: Token = self.consume(TokenType.IDENTIFIER, "Expected subject name")
@@ -358,4 +417,10 @@ class MidasParser(Parser):
self.consume(TokenType.RIGHT_PAREN, "Expected ')' after predicate subject") self.consume(TokenType.RIGHT_PAREN, "Expected ')' after predicate subject")
self.consume(TokenType.EQUAL, "Expected '=' after predicate subject") self.consume(TokenType.EQUAL, "Expected '=' after predicate subject")
condition: Expr = self.constraint() condition: Expr = self.constraint()
return PredicateStmt(name=name, subject=subject, type=type, condition=condition) return PredicateStmt(
location=keyword.location_to(self.previous()),
name=name,
subject=subject,
type=type,
condition=condition,
)
+1 -1
View File
@@ -1,6 +1,7 @@
import ast import ast
from typing import Any, Optional from typing import Any, Optional
from midas.ast.location import Location
from midas.ast.python import ( from midas.ast.python import (
BaseType, BaseType,
ConstraintType, ConstraintType,
@@ -8,7 +9,6 @@ from midas.ast.python import (
FrameType, FrameType,
Function, Function,
FunctionArgument, FunctionArgument,
Location,
MidasType, MidasType,
) )