2 Commits
Author SHA1 Message Date
HEL 4a3363a3d6 feat(checker): add cast expression 2026-05-29 22:04:03 +02:00
HEL 0a3216e07d feat(parser):add cast expression 2026-05-29 22:03:39 +02:00
10 changed files with 73 additions and 47 deletions
@@ -3,6 +3,6 @@
midas.using("02_simple_types.midas")
distance: Meter = 123.45
time: Second = 6.7
distance: Meter = cast(Meter, 123.45)
time: Second = cast(Second, 6.7)
speed = distance / time
+5
View File
@@ -128,4 +128,9 @@ class SetExpr:
value: Expr
class CastExpr:
type: MidasType
expr: Expr
###<
+10
View File
@@ -555,3 +555,13 @@ class PythonAstPrinter(
self._write_line("value", last=True)
with self._child_level(single=True):
expr.value.accept(self)
def visit_cast_expr(self, expr: p.CastExpr) -> None:
self._write_line("CastExpr")
with self._child_level():
self._write_line("type")
with self._child_level(single=True):
expr.type.accept(self)
self._write_line("expr", last=True)
with self._child_level(single=True):
expr.expr.accept(self)
+12
View File
@@ -204,6 +204,9 @@ class Expr(ABC):
@abstractmethod
def visit_set_expr(self, expr: SetExpr) -> T: ...
@abstractmethod
def visit_cast_expr(self, expr: CastExpr) -> T: ...
@dataclass(frozen=True)
class BinaryExpr(Expr):
@@ -287,3 +290,12 @@ class SetExpr(Expr):
def accept(self, visitor: Expr.Visitor[T]) -> T:
return visitor.visit_set_expr(self)
@dataclass(frozen=True)
class CastExpr(Expr):
type: MidasType
expr: Expr
def accept(self, visitor: Expr.Visitor[T]) -> T:
return visitor.visit_cast_expr(self)
+3
View File
@@ -291,6 +291,9 @@ class Checker(
def visit_set_expr(self, expr: p.SetExpr) -> Type: ...
def visit_cast_expr(self, expr: p.CastExpr) -> Type:
return expr.type.accept(self)
def visit_base_type(self, node: p.BaseType) -> Type:
return self.ctx.get_type(node.base)
+27 -13
View File
@@ -7,6 +7,7 @@ from midas.ast.python import (
BaseType,
BinaryExpr,
CallExpr,
CastExpr,
CompareExpr,
ConstraintType,
Expr,
@@ -38,6 +39,8 @@ class UnsupportedSyntaxError(Exception):
class PythonParser:
CAST_FUNCTION = "cast"
def parse_module(self, node: ast.Module) -> list[Stmt]:
statements: list[Stmt] = []
for stmt in node.body:
@@ -90,15 +93,14 @@ class PythonParser:
value=value,
simple=1,
):
type = self._parse_type(annotation, root=True)
if type is not None:
statements.append(
TypeAssign(
location=loc,
name=target,
type=type,
)
type = self._parse_type(annotation)
statements.append(
TypeAssign(
location=loc,
name=target,
type=type,
)
)
if value is not None:
statements.append(
@@ -215,9 +217,7 @@ class PythonParser:
default=default,
)
def _parse_type(
self, type_expr: ast.expr, root: bool = False
) -> Optional[MidasType]:
def _parse_type(self, type_expr: ast.expr) -> MidasType:
loc: Location = Location.from_ast(type_expr)
match type_expr:
case ast.Subscript(value=ast.Name(id="Frame"), slice=schema):
@@ -265,8 +265,6 @@ class PythonParser:
)
case _:
if root:
return None
raise UnsupportedSyntaxError(type_expr)
def _parse_frame_type(self, schema: ast.expr) -> FrameType:
@@ -339,6 +337,9 @@ class PythonParser:
case ast.Compare():
return self.parse_compare(node)
case ast.Call(func=ast.Name(id=self.CAST_FUNCTION)):
return self.parse_cast(node)
case ast.Call():
return self.parse_call(node)
@@ -407,6 +408,19 @@ class PythonParser:
)
return expr
def parse_cast(self, node: ast.Call) -> CastExpr:
match node:
case ast.Call(args=[type, expr], keywords=[]):
return CastExpr(
location=Location.from_ast(node),
type=self._parse_type(type),
expr=self.parse_expr(expr),
)
case _:
raise InvalidSyntaxError(
f"Invalid call to {self.CAST_FUNCTION}, expected type and expression"
)
def parse_call(self, node: ast.Call) -> CallExpr:
return CallExpr(
location=Location.from_ast(node),
+3
View File
@@ -114,3 +114,6 @@ class Resolver(p.Stmt.Visitor[None], p.Expr.Visitor[None]):
def visit_set_expr(self, expr: p.SetExpr) -> None:
self.resolve(expr.value)
self.resolve(expr.object)
def visit_cast_expr(self, expr: p.CastExpr) -> None:
self.resolve(expr.expr)
+2 -2
View File
@@ -3,6 +3,6 @@
midas.using("04_custom_types.midas")
distance: Meter = 123.45
time: Second = 6.7
distance: Meter = cast(Meter, 123.45)
time: Second = cast(Second, 6.7)
speed = distance / time
@@ -1,32 +1,3 @@
{
"diagnostics": [
{
"type": "Error",
"location": {
"start": [
6,
0
],
"end": [
6,
24
]
},
"message": "Cannot assign BaseType(name='float') to distance of type SimpleType(name='Meter', base=BaseType(name='float'))"
},
{
"type": "Error",
"location": {
"start": [
7,
0
],
"end": [
7,
18
]
},
"message": "Cannot assign BaseType(name='float') to time of type SimpleType(name='Second', base=BaseType(name='float'))"
}
]
"diagnostics": []
}
+8
View File
@@ -6,6 +6,7 @@ from midas.ast.python import (
BaseType,
BinaryExpr,
CallExpr,
CastExpr,
CompareExpr,
ConstraintType,
Expr,
@@ -228,3 +229,10 @@ class PythonAstJsonSerializer(
"name": expr.name,
"value": expr.value.accept(self),
}
def visit_cast_expr(self, expr: CastExpr) -> dict:
return {
"_type": "CastExpr",
"type": expr.type.accept(self),
"expr": expr.expr.accept(self),
}