Compare commits
15
Commits
c6ead886ec
...
a1f2937e16
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a1f2937e16
|
||
|
|
2063d94dce
|
||
|
|
22fc8010d8
|
||
|
|
aff1097d91
|
||
|
|
12d034fd1e
|
||
|
|
200709cca6
|
||
|
|
700284296c
|
||
|
|
0b53259b90
|
||
|
|
0461a4184c
|
||
|
|
01d6e41893
|
||
|
|
80e611e49c
|
||
|
|
c00915966f
|
||
|
|
beaa4d95d8
|
||
|
|
bfa0bb3ee0
|
||
|
|
31158df2a9
|
+16
-4
@@ -4,6 +4,7 @@
|
||||
###> Imports
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum, auto
|
||||
from typing import Any, Generic, Optional, TypeVar
|
||||
|
||||
from midas.ast.location import Location
|
||||
@@ -20,6 +21,11 @@ class TypeParam:
|
||||
bound: Optional[Type]
|
||||
|
||||
|
||||
class MemberKind(Enum):
|
||||
PROPERTY = auto()
|
||||
METHOD = auto()
|
||||
|
||||
|
||||
###<
|
||||
|
||||
|
||||
@@ -30,15 +36,16 @@ class TypeStmt:
|
||||
type: Type
|
||||
|
||||
|
||||
class PropertyStmt:
|
||||
class MemberStmt:
|
||||
name: Token
|
||||
type: Type
|
||||
kind: MemberKind
|
||||
|
||||
|
||||
class ExtendStmt:
|
||||
name: Token
|
||||
params: list[TypeParam]
|
||||
type: Type
|
||||
operations: list[OpStmt]
|
||||
members: list[MemberStmt]
|
||||
|
||||
|
||||
class OpStmt:
|
||||
@@ -118,7 +125,12 @@ class ConstraintType:
|
||||
|
||||
|
||||
class ComplexType:
|
||||
properties: list[PropertyStmt]
|
||||
members: list[MemberStmt]
|
||||
|
||||
|
||||
class ExtensionType:
|
||||
base: Type
|
||||
extension: ComplexType
|
||||
|
||||
|
||||
class FunctionType:
|
||||
|
||||
+25
-6
@@ -7,6 +7,7 @@ from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum, auto
|
||||
from typing import Any, Generic, Optional, TypeVar
|
||||
|
||||
from midas.ast.location import Location
|
||||
@@ -21,6 +22,11 @@ class TypeParam:
|
||||
bound: Optional[Type]
|
||||
|
||||
|
||||
class MemberKind(Enum):
|
||||
PROPERTY = auto()
|
||||
METHOD = auto()
|
||||
|
||||
|
||||
##############
|
||||
# Statements #
|
||||
##############
|
||||
@@ -38,7 +44,7 @@ class Stmt(ABC):
|
||||
def visit_type_stmt(self, stmt: TypeStmt) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_property_stmt(self, stmt: PropertyStmt) -> T: ...
|
||||
def visit_member_stmt(self, stmt: MemberStmt) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_extend_stmt(self, stmt: ExtendStmt) -> T: ...
|
||||
@@ -61,19 +67,20 @@ class TypeStmt(Stmt):
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PropertyStmt(Stmt):
|
||||
class MemberStmt(Stmt):
|
||||
name: Token
|
||||
type: Type
|
||||
kind: MemberKind
|
||||
|
||||
def accept(self, visitor: Stmt.Visitor[T]) -> T:
|
||||
return visitor.visit_property_stmt(self)
|
||||
return visitor.visit_member_stmt(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ExtendStmt(Stmt):
|
||||
name: Token
|
||||
params: list[TypeParam]
|
||||
type: Type
|
||||
operations: list[OpStmt]
|
||||
members: list[MemberStmt]
|
||||
|
||||
def accept(self, visitor: Stmt.Visitor[T]) -> T:
|
||||
return visitor.visit_extend_stmt(self)
|
||||
@@ -233,6 +240,9 @@ class Type(ABC):
|
||||
@abstractmethod
|
||||
def visit_complex_type(self, type: ComplexType) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_extension_type(self, type: ExtensionType) -> T: ...
|
||||
|
||||
@abstractmethod
|
||||
def visit_function_type(self, type: FunctionType) -> T: ...
|
||||
|
||||
@@ -265,12 +275,21 @@ class ConstraintType(Type):
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ComplexType(Type):
|
||||
properties: list[PropertyStmt]
|
||||
members: list[MemberStmt]
|
||||
|
||||
def accept(self, visitor: Type.Visitor[T]) -> T:
|
||||
return visitor.visit_complex_type(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ExtensionType(Type):
|
||||
base: Type
|
||||
extension: ComplexType
|
||||
|
||||
def accept(self, visitor: Type.Visitor[T]) -> T:
|
||||
return visitor.visit_extension_type(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FunctionType(Type):
|
||||
pos_args: list[Argument]
|
||||
|
||||
+49
-22
@@ -111,9 +111,10 @@ class MidasAstPrinter(
|
||||
self._write_line(f'name: "{param.name.lexeme}"')
|
||||
self._write_optional_child("bound", param.bound, last=True)
|
||||
|
||||
def visit_property_stmt(self, stmt: m.PropertyStmt):
|
||||
self._write_line("PropertyStmt")
|
||||
def visit_member_stmt(self, stmt: m.MemberStmt):
|
||||
self._write_line("MemberStmt")
|
||||
with self._child_level():
|
||||
self._write_line(f"kind: {stmt.kind.name}")
|
||||
self._write_line(f'name: "{stmt.name.lexeme}"')
|
||||
self._write_line("type", last=True)
|
||||
with self._child_level(single=True):
|
||||
@@ -129,16 +130,21 @@ class MidasAstPrinter(
|
||||
if i == len(stmt.params) - 1:
|
||||
self._mark_last()
|
||||
self._print_type_param(param)
|
||||
self._write_line("type")
|
||||
with self._child_level(single=True):
|
||||
stmt.type.accept(self)
|
||||
self._write_line("operations", last=True)
|
||||
self._write_line(f'name: "{stmt.name.lexeme}"')
|
||||
self._write_line("params")
|
||||
with self._child_level():
|
||||
for i, op in enumerate(stmt.operations):
|
||||
for i, param in enumerate(stmt.params):
|
||||
self._idx = i
|
||||
if i == len(stmt.operations) - 1:
|
||||
if i == len(stmt.params) - 1:
|
||||
self._mark_last()
|
||||
op.accept(self)
|
||||
self._print_type_param(param)
|
||||
self._write_line("members", last=True)
|
||||
with self._child_level():
|
||||
for i, member in enumerate(stmt.members):
|
||||
self._idx = i
|
||||
if i == len(stmt.members) - 1:
|
||||
self._mark_last()
|
||||
member.accept(self)
|
||||
|
||||
def visit_op_stmt(self, stmt: m.OpStmt) -> None:
|
||||
self._write_line("OpStmt")
|
||||
@@ -262,13 +268,23 @@ class MidasAstPrinter(
|
||||
def visit_complex_type(self, type: m.ComplexType) -> None:
|
||||
self._write_line("ComplexType")
|
||||
with self._child_level():
|
||||
self._write_line("properties", last=True)
|
||||
self._write_line("members", last=True)
|
||||
with self._child_level():
|
||||
for i, prop in enumerate(type.properties):
|
||||
for i, member in enumerate(type.members):
|
||||
self._idx = i
|
||||
if i == len(type.properties) - 1:
|
||||
if i == len(type.members) - 1:
|
||||
self._mark_last()
|
||||
prop.accept(self)
|
||||
member.accept(self)
|
||||
|
||||
def visit_extension_type(self, type: m.ExtensionType) -> None:
|
||||
self._write_line("ExtensionType")
|
||||
with self._child_level():
|
||||
self._write_line("base")
|
||||
with self._child_level(single=True):
|
||||
type.base.accept(self)
|
||||
self._write_line("extension", last=True)
|
||||
with self._child_level(single=True):
|
||||
type.extension.accept(self)
|
||||
|
||||
def visit_function_type(self, type: m.FunctionType) -> None:
|
||||
self._write_line("FunctionType")
|
||||
@@ -332,16 +348,24 @@ class MidasPrinter(m.Expr.Visitor[str], m.Stmt.Visitor[str], m.Type.Visitor[str]
|
||||
res += "<:" + param.bound.accept(self)
|
||||
return res
|
||||
|
||||
def visit_property_stmt(self, stmt: m.PropertyStmt):
|
||||
res: str = f"{stmt.name.lexeme}: {stmt.type.accept(self)}"
|
||||
def visit_member_stmt(self, stmt: m.MemberStmt):
|
||||
keyword: str = {
|
||||
m.MemberKind.PROPERTY: "prop",
|
||||
m.MemberKind.METHOD: "def",
|
||||
}.get(stmt.kind, "")
|
||||
res: str = f"{keyword} {stmt.name.lexeme}: {stmt.type.accept(self)}"
|
||||
return self.indented(res)
|
||||
|
||||
def visit_extend_stmt(self, stmt: m.ExtendStmt):
|
||||
res: str = self.indented(f"extend {stmt.type.accept(self)}")
|
||||
template: str = ""
|
||||
if len(stmt.params) != 0:
|
||||
params: list[str] = [self._print_type_param(param) for param in stmt.params]
|
||||
template = f"[{', '.join(params)}]"
|
||||
res: str = self.indented(f"extend {stmt.name.lexeme}{template}")
|
||||
res += " {\n"
|
||||
self.level += 1
|
||||
for op in stmt.operations:
|
||||
res += op.accept(self)
|
||||
for member in stmt.members:
|
||||
res += member.accept(self) + "\n"
|
||||
self.level -= 1
|
||||
res += self.indented("}")
|
||||
return res
|
||||
@@ -411,16 +435,19 @@ class MidasPrinter(m.Expr.Visitor[str], m.Stmt.Visitor[str], m.Type.Visitor[str]
|
||||
def visit_complex_type(self, type: m.ComplexType) -> str:
|
||||
res: str = "{\n"
|
||||
self.level += 1
|
||||
for prop in type.properties:
|
||||
res += prop.accept(self)
|
||||
for member in type.members:
|
||||
res += member.accept(self)
|
||||
res += "\n"
|
||||
self.level -= 1
|
||||
res += self.indented("}")
|
||||
return res
|
||||
|
||||
def visit_extension_type(self, type: m.ExtensionType) -> str:
|
||||
return f"{type.base.accept(self)} & {type.extension.accept(self)}"
|
||||
|
||||
def visit_function_type(self, type: m.FunctionType) -> str:
|
||||
pos_args: list[str] = [self._print_arg(arg) for arg in type.pos_args]
|
||||
kw_args: list[str] = [self._print_arg(arg) for arg in type.pos_args]
|
||||
kw_args: list[str] = [self._print_arg(arg) for arg in type.kw_args]
|
||||
args: list[str] = pos_args
|
||||
|
||||
if len(pos_args) != 0:
|
||||
@@ -429,7 +456,7 @@ class MidasPrinter(m.Expr.Visitor[str], m.Stmt.Visitor[str], m.Type.Visitor[str]
|
||||
args.append("*")
|
||||
args += kw_args
|
||||
|
||||
return f"({', '.join(args)}) -> {type.returns.accept(self)}"
|
||||
return f"fn ({', '.join(args)}) -> {type.returns.accept(self)}"
|
||||
|
||||
def _print_arg(self, arg: m.FunctionType.Argument) -> str:
|
||||
res: str = ""
|
||||
|
||||
+43
-17
@@ -8,6 +8,7 @@ from midas.checker.reporter import FileReporter, Reporter
|
||||
from midas.checker.types import (
|
||||
AliasType,
|
||||
ComplexType,
|
||||
ExtensionType,
|
||||
Function,
|
||||
GenericType,
|
||||
Type,
|
||||
@@ -29,6 +30,8 @@ class MidasTyper(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[Type
|
||||
self.types: TypesRegistry = types
|
||||
self._local_variables: dict[str, TypeVar] = {}
|
||||
|
||||
self._current_name: Optional[str] = None
|
||||
|
||||
define_builtins(self.types)
|
||||
|
||||
def process(self, source: str, path: Optional[str]):
|
||||
@@ -65,9 +68,10 @@ class MidasTyper(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[Type
|
||||
stmt.accept(self)
|
||||
|
||||
def visit_type_stmt(self, stmt: m.TypeStmt) -> None:
|
||||
name: str = stmt.name.lexeme
|
||||
self._current_name = name
|
||||
params: list[TypeVar] = self._resolve_type_params(stmt.params)
|
||||
|
||||
name: str = stmt.name.lexeme
|
||||
type: Type = stmt.type.accept(self)
|
||||
if len(params) != 0:
|
||||
type = GenericType(name=name, params=params, body=type)
|
||||
@@ -75,20 +79,25 @@ class MidasTyper(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[Type
|
||||
type = AliasType(name=name, type=type)
|
||||
self.types.define_type(name, type)
|
||||
self._local_variables.clear()
|
||||
self._current_name = None
|
||||
|
||||
def visit_property_stmt(self, stmt: m.PropertyStmt) -> None: ...
|
||||
def visit_member_stmt(self, stmt: m.MemberStmt) -> None: ...
|
||||
|
||||
def visit_extend_stmt(self, stmt: m.ExtendStmt) -> None:
|
||||
self._resolve_type_params(stmt.params)
|
||||
base: Type = stmt.type.accept(self)
|
||||
for op in stmt.operations:
|
||||
right: Type = op.operand.accept(self)
|
||||
result: Type = op.result.accept(self)
|
||||
self.types.define_operation(
|
||||
left=base,
|
||||
operator=op.name.lexeme,
|
||||
right=right,
|
||||
result=result,
|
||||
base_name: str = stmt.name.lexeme
|
||||
try:
|
||||
_ = self.get_type(base_name)
|
||||
except NameError:
|
||||
self.reporter.error(stmt.name.get_location(), f"Unknown type '{base_name}'")
|
||||
|
||||
for member in stmt.members:
|
||||
member_type: Type = member.type.accept(self)
|
||||
self.types.define_member(
|
||||
base_name,
|
||||
member.name.lexeme,
|
||||
member_type,
|
||||
member.kind == m.MemberKind.METHOD,
|
||||
)
|
||||
|
||||
def visit_op_stmt(self, stmt: m.OpStmt) -> None: ...
|
||||
@@ -113,12 +122,24 @@ class MidasTyper(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[Type
|
||||
def visit_wildcard_expr(self, expr: m.WildcardExpr) -> None: ...
|
||||
|
||||
def visit_named_type(self, type: m.NamedType) -> Type:
|
||||
return self.get_type(type.name.lexeme)
|
||||
name: str = type.name.lexeme
|
||||
try:
|
||||
return self.get_type(name)
|
||||
except NameError:
|
||||
msg: str = f"Undefined type {name}"
|
||||
if self._current_name == name:
|
||||
msg += ". Recursive types are not supported, use an extend block"
|
||||
self.reporter.error(type.name.get_location(), msg)
|
||||
return UnknownType()
|
||||
|
||||
def visit_generic_type(self, type: m.GenericType) -> Type:
|
||||
type_: Type = type.type.accept(self)
|
||||
args: list[Type] = [arg.accept(self) for arg in type.args]
|
||||
return self.types.apply_generic(type_, args)
|
||||
try:
|
||||
return self.types.apply_generic(type_, args)
|
||||
except Exception as e:
|
||||
self.reporter.error(type.location, f"Cannot apply generic type: {e}")
|
||||
return UnknownType()
|
||||
|
||||
def visit_constraint_type(self, type: m.ConstraintType) -> Type:
|
||||
type_: Type = type.type.accept(self)
|
||||
@@ -126,16 +147,21 @@ class MidasTyper(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[Type
|
||||
# TODO
|
||||
return UnknownType()
|
||||
|
||||
def visit_complex_type(self, type: m.ComplexType) -> Type:
|
||||
def visit_complex_type(self, type: m.ComplexType) -> ComplexType:
|
||||
return ComplexType(
|
||||
properties={
|
||||
prop.name.lexeme: prop.type.accept(self) for prop in type.properties
|
||||
members={
|
||||
member.name.lexeme: member.type.accept(self) for member in type.members
|
||||
}
|
||||
)
|
||||
|
||||
def visit_extension_type(self, type: m.ExtensionType) -> Type:
|
||||
return ExtensionType(
|
||||
base=type.base.accept(self),
|
||||
extension=self.visit_complex_type(type.extension),
|
||||
)
|
||||
|
||||
def visit_function_type(self, type: m.FunctionType) -> Type:
|
||||
return Function(
|
||||
name="<anonymous>",
|
||||
pos_args=[
|
||||
Function.Argument(
|
||||
pos=i,
|
||||
|
||||
+61
-87
@@ -11,13 +11,10 @@ from midas.checker.registry import TypesRegistry
|
||||
from midas.checker.reporter import FileReporter, Reporter
|
||||
from midas.checker.resolver import Resolver
|
||||
from midas.checker.types import (
|
||||
ComplexType,
|
||||
Function,
|
||||
Operation,
|
||||
Type,
|
||||
UnitType,
|
||||
UnknownType,
|
||||
unfold_type,
|
||||
)
|
||||
from midas.parser.python import PythonParser
|
||||
|
||||
@@ -192,7 +189,6 @@ class PythonTyper(
|
||||
returns_hint = stmt.returns.accept(self)
|
||||
# Early define to handle simple fully-typed recursion
|
||||
inside_function: Function = Function(
|
||||
name=stmt.name,
|
||||
pos_args=pos_args,
|
||||
args=args,
|
||||
kw_args=kw_args,
|
||||
@@ -227,7 +223,6 @@ class PythonTyper(
|
||||
|
||||
# TODO: handle *args and **kwargs sinks
|
||||
function: Function = Function(
|
||||
name=stmt.name,
|
||||
pos_args=pos_args,
|
||||
args=args,
|
||||
kw_args=kw_args,
|
||||
@@ -250,8 +245,8 @@ class PythonTyper(
|
||||
case p.VariableExpr():
|
||||
self._assign_var(location, target, value_type)
|
||||
|
||||
case p.GetExpr():
|
||||
self._assign_attr(location, target, value_type)
|
||||
case p.GetExpr(object=object, name=name):
|
||||
self._assign_attr(location, object, name, value_type)
|
||||
|
||||
case _:
|
||||
if not isinstance(target, p.VariableExpr):
|
||||
@@ -276,33 +271,20 @@ class PythonTyper(
|
||||
f"Cannot assign {value_type} to variable '{name}' of type {var_type}",
|
||||
)
|
||||
|
||||
def _assign_attr(self, location: Location, target: p.GetExpr, value_type: Type):
|
||||
object: Type = self.type_of(target.object)
|
||||
base_object: Type = unfold_type(object)
|
||||
match base_object:
|
||||
case ComplexType(properties=properties):
|
||||
if target.name not in properties:
|
||||
self.reporter.error(
|
||||
target.location, f"Unknown property '{object}.{target.name}'"
|
||||
)
|
||||
return
|
||||
|
||||
prop_type: Type = properties[target.name]
|
||||
if not self.is_subtype(value_type, prop_type):
|
||||
self.reporter.error(
|
||||
location,
|
||||
f"Cannot assign {value_type} to property '{object}.{target.name}' of type {prop_type}",
|
||||
)
|
||||
return
|
||||
|
||||
case UnknownType():
|
||||
pass
|
||||
|
||||
case _:
|
||||
self.reporter.error(
|
||||
target.location,
|
||||
f"Cannot assign {value_type} to unknown property '{object}.{target.name}'",
|
||||
)
|
||||
def _assign_attr(
|
||||
self, location: Location, object: p.Expr, name: str, value_type: Type
|
||||
):
|
||||
object_type: Type = self.type_of(object)
|
||||
member: Optional[Type] = self.types.lookup_member(object_type, name)
|
||||
if member is None:
|
||||
self.reporter.error(location, f"Unknown member '{name}' of {object_type}")
|
||||
return
|
||||
self.logger.debug(f"Member '{name}' of {object_type} has type {member}")
|
||||
if not self.is_subtype(value_type, member):
|
||||
self.reporter.error(
|
||||
location,
|
||||
f"Cannot assign {value_type} to member '{object_type}.{name}' of type {member}",
|
||||
)
|
||||
|
||||
def visit_return_stmt(self, stmt: p.ReturnStmt) -> None:
|
||||
type: Type = stmt.value.accept(self) if stmt.value is not None else UnitType()
|
||||
@@ -341,47 +323,36 @@ class PythonTyper(
|
||||
left: Type = self.type_of(expr.left)
|
||||
right: Type = self.type_of(expr.right)
|
||||
|
||||
operations: list[Operation] = self.types.get_operations_by_name(method)
|
||||
valid_operations: list[Operation] = []
|
||||
for op in operations:
|
||||
sig: Operation.CallSignature = op.signature
|
||||
if self.is_subtype(left, sig.left) and self.is_subtype(right, sig.right):
|
||||
valid_operations.append(op)
|
||||
|
||||
if len(valid_operations) == 0:
|
||||
operation: Optional[Type] = self.types.lookup_member(left, method)
|
||||
if operation is None:
|
||||
self.reporter.error(
|
||||
expr.location,
|
||||
f"Undefined operation {method} between {left} and {right}",
|
||||
)
|
||||
return UnknownType()
|
||||
elif len(valid_operations) == 1:
|
||||
self.logger.debug(f"Unique operation {method} between {left} and {right}")
|
||||
return valid_operations[0].result
|
||||
|
||||
for i, op1 in enumerate(valid_operations):
|
||||
sig1: Operation.CallSignature = op1.signature
|
||||
best_match: bool = True
|
||||
for j, op2 in enumerate(valid_operations):
|
||||
if i == j:
|
||||
continue
|
||||
sig2: Operation.CallSignature = op2.signature
|
||||
match operation:
|
||||
case Function() as function:
|
||||
if not self._is_binary_function(function):
|
||||
self.reporter.error(
|
||||
expr.location,
|
||||
f"Wrong definition of binary operation. Expected function with 2 positional-only parameters, got {function}",
|
||||
)
|
||||
return UnknownType()
|
||||
|
||||
# If op1 is not a full overload of op2 (i.e. operands of op1 are subtypes of op2's)
|
||||
# ambiguity -> not best match
|
||||
if not self.is_subtype(sig1.left, sig2.left) or not self.is_subtype(
|
||||
sig1.right, sig2.right
|
||||
):
|
||||
best_match = False
|
||||
break
|
||||
self.logger.debug(f"{op1} is a full overload of {op2}")
|
||||
if best_match:
|
||||
return op1.result
|
||||
|
||||
self.reporter.error(
|
||||
expr.location,
|
||||
f"Ambiguous operation {method} between {left} and {right}, multiple matching overloads: {', '.join(map(str, valid_operations))}",
|
||||
)
|
||||
return UnknownType()
|
||||
rhs: Function.Argument = function.pos_args[0]
|
||||
if not self.is_subtype(right, rhs.type):
|
||||
self.reporter.error(
|
||||
expr.location,
|
||||
f"Wrong type for right-hand side, expected {rhs.type}, got {right}",
|
||||
)
|
||||
return UnknownType()
|
||||
return function.returns
|
||||
case _:
|
||||
self.reporter.warning(
|
||||
expr.location, f"Unsupported operation {operation}"
|
||||
)
|
||||
return UnknownType()
|
||||
|
||||
def visit_compare_expr(self, expr: p.CompareExpr) -> Type:
|
||||
method: Optional[str] = COMPARATOR_METHODS.get(expr.operator.__class__)
|
||||
@@ -422,24 +393,14 @@ class PythonTyper(
|
||||
|
||||
def visit_get_expr(self, expr: p.GetExpr) -> Type:
|
||||
object: Type = self.type_of(expr.object)
|
||||
base_object: Type = unfold_type(object)
|
||||
match base_object:
|
||||
case ComplexType(properties=properties):
|
||||
if expr.name not in properties:
|
||||
self.reporter.error(
|
||||
expr.location, f"Unknown property '{expr.name} on {object}"
|
||||
)
|
||||
return UnknownType()
|
||||
return properties[expr.name]
|
||||
|
||||
case UnknownType():
|
||||
return UnknownType()
|
||||
|
||||
case _:
|
||||
self.reporter.error(
|
||||
expr.location, f"Cannot get property '{expr.name}' on {object}"
|
||||
)
|
||||
return UnknownType()
|
||||
member: Optional[Type] = self.types.lookup_member(object, expr.name)
|
||||
if member is None:
|
||||
self.reporter.error(
|
||||
expr.location, f"Unknown member '{expr.name}' of {object}"
|
||||
)
|
||||
return UnknownType()
|
||||
self.logger.debug(f"Member '{expr.name}' of {object} has type {member}")
|
||||
return member
|
||||
|
||||
def visit_literal_expr(self, expr: p.LiteralExpr) -> Type:
|
||||
match expr.value:
|
||||
@@ -456,7 +417,11 @@ class PythonTyper(
|
||||
return UnknownType()
|
||||
|
||||
def visit_variable_expr(self, expr: p.VariableExpr) -> Type:
|
||||
return self.look_up_variable(expr.name, expr) or UnknownType()
|
||||
type: Optional[Type] = self.look_up_variable(expr.name, expr)
|
||||
if type is None:
|
||||
self.logger.debug(f"Unknown variable {expr.name} in {self.env.flat_dict()}")
|
||||
self.reporter.warning(expr.location, "Unknown variable")
|
||||
return type or UnknownType()
|
||||
|
||||
def visit_logical_expr(self, expr: p.LogicalExpr) -> Type:
|
||||
left: Type = expr.left.accept(self)
|
||||
@@ -644,3 +609,12 @@ class PythonTyper(
|
||||
)
|
||||
|
||||
return mapped
|
||||
|
||||
def _is_binary_function(self, function: Function) -> bool:
|
||||
if len(function.pos_args) != 1:
|
||||
return False
|
||||
if len(function.args) != 0:
|
||||
return False
|
||||
if len(function.kw_args) != 0:
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from midas.checker.builtins import BUILTIN_SUBTYPES
|
||||
@@ -6,17 +7,23 @@ from midas.checker.types import (
|
||||
AppliedType,
|
||||
BaseType,
|
||||
ComplexType,
|
||||
ExtensionType,
|
||||
Function,
|
||||
GenericType,
|
||||
Operation,
|
||||
OverloadedFunction,
|
||||
Type,
|
||||
TypeVar,
|
||||
UnknownType,
|
||||
substitute_typevars,
|
||||
)
|
||||
|
||||
|
||||
class TypesRegistry:
|
||||
def __init__(self) -> None:
|
||||
self.logger: logging.Logger = logging.getLogger("TypesRegistry")
|
||||
self._types: dict[str, Type] = {}
|
||||
self._members: dict[str, dict[str, Type]] = {}
|
||||
self._operations: dict[Operation.CallSignature, Type] = {}
|
||||
|
||||
def get_type(self, name: str) -> Type:
|
||||
@@ -86,6 +93,28 @@ class TypesRegistry:
|
||||
self._types[name] = type
|
||||
return type
|
||||
|
||||
def define_member(
|
||||
self, type_name: str, member_name: str, member_type: Type, is_method: bool
|
||||
):
|
||||
members: dict[str, Type] = self._members.setdefault(type_name, {})
|
||||
if member_name in members:
|
||||
if not is_method:
|
||||
self.logger.error(
|
||||
f"Member '{member_name}' already defined for type {type_name}"
|
||||
)
|
||||
return
|
||||
current: Type = members[member_name]
|
||||
combined: Type
|
||||
match current:
|
||||
case OverloadedFunction(overloads=overloads):
|
||||
combined = OverloadedFunction(overloads=overloads + [member_type])
|
||||
case _:
|
||||
combined = OverloadedFunction(overloads=[current, member_type])
|
||||
members[member_name] = combined
|
||||
|
||||
else:
|
||||
members[member_name] = member_type
|
||||
|
||||
def define_operation(self, left: Type, operator: str, right: Type, result: Type):
|
||||
"""Define an operation in the registry
|
||||
|
||||
@@ -143,6 +172,11 @@ class TypesRegistry:
|
||||
case (Function(), Function()):
|
||||
return self.is_func_subtype(type1, type2)
|
||||
|
||||
case (TypeVar(bound=bound), _):
|
||||
if bound is None:
|
||||
return False
|
||||
return self.is_subtype(bound, type2)
|
||||
|
||||
return False
|
||||
|
||||
# TODO: verify the logic in here
|
||||
@@ -311,3 +345,51 @@ class TypesRegistry:
|
||||
reduced = True
|
||||
break
|
||||
return [types[i] for i in keep]
|
||||
|
||||
def lookup_member(self, type: Type, member_name: str) -> Optional[Type]:
|
||||
match type:
|
||||
case AliasType(name=name, type=base):
|
||||
if name in self._members:
|
||||
if member_name in self._members[name]:
|
||||
return self._members[name][member_name]
|
||||
return self.lookup_member(base, member_name)
|
||||
|
||||
case AppliedType(name=name, body=body, args=args):
|
||||
generic: Type = self.get_type(name)
|
||||
|
||||
if not isinstance(generic, GenericType):
|
||||
raise ValueError("AppliedType not derived from a GenericType")
|
||||
|
||||
substitutions = {
|
||||
type_var.name: arg for arg, type_var in zip(args, generic.params)
|
||||
}
|
||||
if name in self._members:
|
||||
if member_name in self._members[name]:
|
||||
member_type: Type = self._members[name][member_name]
|
||||
return substitute_typevars(member_type, substitutions)
|
||||
|
||||
member_type2: Optional[Type] = self.lookup_member(body, member_name)
|
||||
if member_type2 is not None:
|
||||
member_type2 = substitute_typevars(member_type2, substitutions)
|
||||
return member_type2
|
||||
|
||||
case ComplexType(members=members):
|
||||
if member_name in members:
|
||||
return members[member_name]
|
||||
self.logger.debug(f"No member '{member_name}' in {type}")
|
||||
return None
|
||||
|
||||
case ExtensionType(base=base, extension=ComplexType(members=members)):
|
||||
if member_name in members:
|
||||
return members[member_name]
|
||||
self.logger.debug(
|
||||
f"No member '{member_name}' on {type}, looking up in base"
|
||||
)
|
||||
return self.lookup_member(base, member_name)
|
||||
|
||||
case UnknownType():
|
||||
return UnknownType()
|
||||
|
||||
case _:
|
||||
self.logger.debug(f"Can't get member on {type}")
|
||||
return None
|
||||
|
||||
+48
-11
@@ -35,7 +35,6 @@ class UnitType:
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class Function:
|
||||
name: str
|
||||
pos_args: list[Argument]
|
||||
args: list[Argument]
|
||||
kw_args: list[Argument]
|
||||
@@ -56,7 +55,7 @@ class Function:
|
||||
args.append("*")
|
||||
args += list(map(str, self.kw_args))
|
||||
|
||||
return f"{self.name}({', '.join(args)}) -> {self.returns}"
|
||||
return f"({', '.join(args)}) -> {self.returns}"
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class Argument:
|
||||
@@ -71,14 +70,31 @@ class Function:
|
||||
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class ComplexType:
|
||||
properties: dict[str, Type]
|
||||
class OverloadedFunction:
|
||||
overloads: list[Type]
|
||||
|
||||
def __str__(self) -> str:
|
||||
props: list[str] = [f"{name}: {type}" for name, type in self.properties.items()]
|
||||
return "<overloaded function>"
|
||||
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class ComplexType:
|
||||
members: dict[str, Type]
|
||||
|
||||
def __str__(self) -> str:
|
||||
props: list[str] = [f"{name}: {type}" for name, type in self.members.items()]
|
||||
return f"{{{', '.join(props)}}}"
|
||||
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class ExtensionType:
|
||||
base: Type
|
||||
extension: ComplexType
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"{self.base} & {self.extension}"
|
||||
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class Operation:
|
||||
signature: CallSignature
|
||||
@@ -141,30 +157,49 @@ def substitute_typevars(type: Type, substitutions: dict[str, Type]) -> Type:
|
||||
case BaseType(name=name) if name in substitutions:
|
||||
return substitutions[name]
|
||||
|
||||
case BaseType():
|
||||
return type
|
||||
|
||||
case AliasType(name=name, type=type2):
|
||||
return AliasType(name=name, type=substitute_typevars(type2, substitutions))
|
||||
|
||||
case Function(
|
||||
name=name,
|
||||
pos_args=pos_args,
|
||||
args=args,
|
||||
kw_args=kw_args,
|
||||
returns=returns,
|
||||
):
|
||||
return Function(
|
||||
name=name,
|
||||
pos_args=list(map(sub_argument, pos_args)),
|
||||
args=list(map(sub_argument, args)),
|
||||
kw_args=list(map(sub_argument, kw_args)),
|
||||
returns=substitute_typevars(returns, substitutions),
|
||||
)
|
||||
|
||||
case ComplexType(properties=properties):
|
||||
properties2: dict[str, Type] = {
|
||||
case ComplexType(members=members):
|
||||
members2: dict[str, Type] = {
|
||||
name: substitute_typevars(prop, substitutions)
|
||||
for name, prop in properties.items()
|
||||
for name, prop in members.items()
|
||||
}
|
||||
return ComplexType(properties=properties2)
|
||||
return ComplexType(members=members2)
|
||||
|
||||
case ExtensionType(base=base, extension=ComplexType(members=members)):
|
||||
return ExtensionType(
|
||||
base=substitute_typevars(base, substitutions),
|
||||
extension=ComplexType(
|
||||
members={
|
||||
name: substitute_typevars(prop, substitutions)
|
||||
for name, prop in members.items()
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
case AppliedType(name=name, args=args, body=body):
|
||||
return AppliedType(
|
||||
name=name,
|
||||
args=[substitute_typevars(arg, substitutions) for arg in args],
|
||||
body=substitute_typevars(body, substitutions),
|
||||
)
|
||||
|
||||
case TypeVar(name=name):
|
||||
if name in substitutions:
|
||||
@@ -192,7 +227,9 @@ Type = (
|
||||
| UnknownType
|
||||
| UnitType
|
||||
| Function
|
||||
| OverloadedFunction
|
||||
| ComplexType
|
||||
| ExtensionType
|
||||
| TypeVar
|
||||
| GenericType
|
||||
| AppliedType
|
||||
|
||||
@@ -232,15 +232,14 @@ class MidasHighlighter(
|
||||
self.wrap(LocatableToken(stmt.name), "type-name")
|
||||
stmt.type.accept(self)
|
||||
|
||||
def visit_property_stmt(self, stmt: m.PropertyStmt) -> None:
|
||||
self.wrap(stmt, "property")
|
||||
def visit_member_stmt(self, stmt: m.MemberStmt) -> None:
|
||||
self.wrap(stmt, "member")
|
||||
stmt.type.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)
|
||||
for member in stmt.members:
|
||||
member.accept(self)
|
||||
|
||||
def visit_op_stmt(self, stmt: m.OpStmt) -> None:
|
||||
self.wrap(stmt, "op")
|
||||
@@ -298,8 +297,8 @@ class MidasHighlighter(
|
||||
|
||||
def visit_complex_type(self, type: m.ComplexType) -> None:
|
||||
self.wrap(type, "complex-type")
|
||||
for prop in type.properties:
|
||||
prop.accept(self)
|
||||
for member in type.members:
|
||||
member.accept(self)
|
||||
|
||||
def visit_function_type(self, type: m.FunctionType) -> None:
|
||||
self.wrap(type, "function")
|
||||
@@ -307,6 +306,11 @@ class MidasHighlighter(
|
||||
arg.type.accept(self)
|
||||
type.returns.accept(self)
|
||||
|
||||
def visit_extension_type(self, type: m.ExtensionType) -> None:
|
||||
self.wrap(type, "extension")
|
||||
type.base.accept(self)
|
||||
type.extension.accept(self)
|
||||
|
||||
|
||||
class DiagnosticsHighlighter(Highlighter):
|
||||
EXTRA_CSS_PATH: Optional[Path] = Path(__file__).parent / "hl_diagnostic.css"
|
||||
|
||||
+30
-4
@@ -88,24 +88,50 @@ def print_diagnostic(lines: list[str], diagnostic: Diagnostic, indent: int = 4):
|
||||
@click.option("-l", "--highlight", type=click.File("w"))
|
||||
@click.option("-t", "--types", type=click.File("r"), multiple=True)
|
||||
@click.option("-v", "--verbose", is_flag=True)
|
||||
@click.option("-j", "--show-judgements", is_flag=True)
|
||||
@click.argument("file", type=click.File("r"))
|
||||
def compile(
|
||||
highlight: Optional[TextIO],
|
||||
types: tuple[TextIO],
|
||||
verbose: bool,
|
||||
show_judgements: bool,
|
||||
file: TextIO,
|
||||
):
|
||||
logging.basicConfig(level=logging.DEBUG if verbose else logging.WARN)
|
||||
source: str = file.read()
|
||||
source_path: Path = Path(file.name).resolve()
|
||||
|
||||
checker = TypeChecker()
|
||||
for path in types:
|
||||
checker.import_midas(Path(path.name).resolve())
|
||||
for types_file in types:
|
||||
checker.import_midas(Path(types_file.name).resolve())
|
||||
|
||||
checker.type_check_source(source, str(Path(file.name).resolve()))
|
||||
diagnostics: list[Diagnostic] = checker.diagnostics
|
||||
checker.type_check_source(source, str(source_path))
|
||||
diagnostics: list[Diagnostic] = checker.diagnostics.copy()
|
||||
lines: list[str] = source.split("\n")
|
||||
files: dict[Optional[str], list[str]] = {None: []}
|
||||
|
||||
if show_judgements:
|
||||
for expr, type in checker.python_typer.judgements:
|
||||
print(f"Judged that {expr} at {expr.location} is of type {type}")
|
||||
diagnostics.append(
|
||||
Diagnostic(
|
||||
file_path=str(source_path),
|
||||
location=expr.location,
|
||||
type=DiagnosticType.INFO,
|
||||
message=f"Type: {type}",
|
||||
)
|
||||
)
|
||||
|
||||
for diagnostic in diagnostics:
|
||||
filename: Optional[str] = diagnostic.file_path
|
||||
if filename is not None and filename not in files:
|
||||
path: Path = Path(filename)
|
||||
if path.exists() and path.is_file():
|
||||
files[filename] = path.read_text().split("\n")
|
||||
else:
|
||||
files[filename] = []
|
||||
|
||||
lines: list[str] = files[filename]
|
||||
print_diagnostic(lines, diagnostic)
|
||||
|
||||
if verbose:
|
||||
|
||||
@@ -50,6 +50,9 @@ class TokenType(Enum):
|
||||
PREDICATE = auto()
|
||||
EXTEND = auto()
|
||||
WHERE = auto()
|
||||
PROP = auto()
|
||||
DEF = auto()
|
||||
FUNC = auto()
|
||||
|
||||
# Misc
|
||||
COMMENT = auto()
|
||||
@@ -67,6 +70,9 @@ KEYWORDS: dict[str, TokenType] = {
|
||||
"true": TokenType.TRUE,
|
||||
"false": TokenType.FALSE,
|
||||
"none": TokenType.NONE,
|
||||
"prop": TokenType.PROP,
|
||||
"def": TokenType.DEF,
|
||||
"fn": TokenType.FUNC,
|
||||
}
|
||||
|
||||
|
||||
|
||||
+51
-25
@@ -7,16 +7,18 @@ from midas.ast.midas import (
|
||||
ConstraintType,
|
||||
Expr,
|
||||
ExtendStmt,
|
||||
ExtensionType,
|
||||
FunctionType,
|
||||
GenericType,
|
||||
GetExpr,
|
||||
GroupingExpr,
|
||||
LiteralExpr,
|
||||
LogicalExpr,
|
||||
MemberKind,
|
||||
MemberStmt,
|
||||
NamedType,
|
||||
OpStmt,
|
||||
PredicateStmt,
|
||||
PropertyStmt,
|
||||
Stmt,
|
||||
Type,
|
||||
TypeParam,
|
||||
@@ -163,7 +165,19 @@ class MidasParser(Parser):
|
||||
Returns:
|
||||
TypeExpr: the parsed type expression
|
||||
"""
|
||||
return self.constraint_type()
|
||||
base: Type
|
||||
if self.match(TokenType.FUNC):
|
||||
base = self.function()
|
||||
else:
|
||||
base = self.constraint_type()
|
||||
if self.match(TokenType.AND):
|
||||
extension: ComplexType = self.complex_type()
|
||||
return ExtensionType(
|
||||
location=Location.span(base.location, extension.location),
|
||||
base=base,
|
||||
extension=extension,
|
||||
)
|
||||
return base
|
||||
|
||||
def constraint_type(self) -> Type:
|
||||
type: Type = self.base_type()
|
||||
@@ -215,30 +229,32 @@ class MidasParser(Parser):
|
||||
name=name,
|
||||
)
|
||||
|
||||
def complex_type(self) -> Type:
|
||||
def complex_type(self) -> ComplexType:
|
||||
"""Parse a type definition body
|
||||
|
||||
A type definition body is a set of whitespace-separated
|
||||
property statements enclosed in curly braces
|
||||
|
||||
Returns:
|
||||
list[PropertyStmt]: the parsed type properties
|
||||
ComplexType: the parsed complex type
|
||||
"""
|
||||
left: Token = self.consume(
|
||||
TokenType.LEFT_BRACE, "Expected '{' to start type body"
|
||||
)
|
||||
properties: list[PropertyStmt] = []
|
||||
members: list[MemberStmt] = []
|
||||
# TODO: add keyword to differentiate properties and methods,
|
||||
# and allow multiple methods with the same name but not properties
|
||||
names: set[str] = set()
|
||||
while not self.check(TokenType.RIGHT_BRACE) and not self.is_at_end():
|
||||
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)
|
||||
member: MemberStmt = self.member_stmt()
|
||||
# if member.name.lexeme in names:
|
||||
# raise self.error(member.name, "Duplicate property")
|
||||
# names.add(member.name.lexeme)
|
||||
members.append(member)
|
||||
right: Token = self.consume(TokenType.RIGHT_BRACE, "Unclosed type body")
|
||||
return ComplexType(
|
||||
location=left.location_to(right),
|
||||
properties=properties,
|
||||
members=members,
|
||||
)
|
||||
|
||||
def constraint(self) -> Expr:
|
||||
@@ -376,21 +392,31 @@ class MidasParser(Parser):
|
||||
return True
|
||||
return False
|
||||
|
||||
def property_stmt(self) -> PropertyStmt:
|
||||
"""Parse a property statement
|
||||
def member_stmt(self) -> MemberStmt:
|
||||
"""Parse a member statement
|
||||
|
||||
A type property statement is written `name: Type` or `name: Type where Condition`
|
||||
A type member statement is written `prop name: Type` or `def name: Type`
|
||||
|
||||
Returns:
|
||||
PropertyStmt: the parsed property statement
|
||||
MemberStmt: the parsed member statement
|
||||
"""
|
||||
name: Token = self.consume_identifier("Expected property name")
|
||||
self.consume(TokenType.COLON, "Expected ':' after property name")
|
||||
kind: MemberKind
|
||||
if self.match(TokenType.PROP):
|
||||
kind = MemberKind.PROPERTY
|
||||
elif self.match(TokenType.DEF):
|
||||
kind = MemberKind.METHOD
|
||||
else:
|
||||
raise self.error(self.peek(), "Expected 'prop' or 'def'")
|
||||
|
||||
name: Token = self.consume_identifier("Expected member name")
|
||||
self.consume(TokenType.COLON, "Expected ':' after member name")
|
||||
|
||||
type: Type = self.type_expr()
|
||||
return PropertyStmt(
|
||||
return MemberStmt(
|
||||
location=name.location_to(self.previous()),
|
||||
name=name,
|
||||
type=type,
|
||||
kind=kind,
|
||||
)
|
||||
|
||||
def extend_declaration(self) -> ExtendStmt:
|
||||
@@ -402,20 +428,20 @@ class MidasParser(Parser):
|
||||
ExtendStmt: the parsed extension statement
|
||||
"""
|
||||
keyword: Token = self.previous()
|
||||
name: Token = self.consume_identifier("Expected type name")
|
||||
params: list[TypeParam] = self.type_params()
|
||||
|
||||
type: Type = self.type_expr()
|
||||
self.consume(TokenType.LEFT_BRACE, "Expected '{' to start extend body")
|
||||
operations: list[OpStmt] = []
|
||||
members: list[MemberStmt] = []
|
||||
while not self.is_at_end() and not self.check(TokenType.RIGHT_BRACE):
|
||||
operations.append(self.op_declaration())
|
||||
members.append(self.member_stmt())
|
||||
self.consume(TokenType.RIGHT_BRACE, "Unclosed extend body")
|
||||
location: Location = keyword.location_to(self.previous())
|
||||
return ExtendStmt(
|
||||
location=location,
|
||||
name=name,
|
||||
params=params,
|
||||
type=type,
|
||||
operations=operations,
|
||||
members=members,
|
||||
)
|
||||
|
||||
def op_declaration(self) -> OpStmt:
|
||||
@@ -487,12 +513,12 @@ class MidasParser(Parser):
|
||||
name = self.advance()
|
||||
self.advance()
|
||||
type: Type = self.type_expr()
|
||||
required: bool = self.match(TokenType.QMARK)
|
||||
optional: bool = self.match(TokenType.QMARK)
|
||||
arg = FunctionType.Argument(
|
||||
location=None,
|
||||
name=name,
|
||||
type=type,
|
||||
required=required,
|
||||
required=not optional,
|
||||
)
|
||||
if positional:
|
||||
pos_args.append(arg)
|
||||
|
||||
@@ -6,16 +6,17 @@ from midas.ast.midas import (
|
||||
ConstraintType,
|
||||
Expr,
|
||||
ExtendStmt,
|
||||
ExtensionType,
|
||||
FunctionType,
|
||||
GenericType,
|
||||
GetExpr,
|
||||
GroupingExpr,
|
||||
LiteralExpr,
|
||||
LogicalExpr,
|
||||
MemberStmt,
|
||||
NamedType,
|
||||
OpStmt,
|
||||
PredicateStmt,
|
||||
PropertyStmt,
|
||||
Stmt,
|
||||
Type,
|
||||
TypeParam,
|
||||
@@ -58,9 +59,10 @@ class MidasAstJsonSerializer(
|
||||
"bound": self._serialize_optional(param.bound),
|
||||
}
|
||||
|
||||
def visit_property_stmt(self, stmt: PropertyStmt) -> dict:
|
||||
def visit_member_stmt(self, stmt: MemberStmt) -> dict:
|
||||
return {
|
||||
"_type": "PropertyStmt",
|
||||
"_type": "MemberStmt",
|
||||
"kind": stmt.kind.name,
|
||||
"name": stmt.name.lexeme,
|
||||
"type": stmt.type.accept(self),
|
||||
}
|
||||
@@ -68,8 +70,9 @@ class MidasAstJsonSerializer(
|
||||
def visit_extend_stmt(self, stmt: ExtendStmt) -> dict:
|
||||
return {
|
||||
"_type": "ExtendStmt",
|
||||
"type": stmt.type.accept(self),
|
||||
"operations": self._serialize_list(stmt.operations),
|
||||
"name": stmt.name.lexeme,
|
||||
"params": [self._serialize_type_param(param) for param in stmt.params],
|
||||
"members": self._serialize_list(stmt.members),
|
||||
}
|
||||
|
||||
def visit_op_stmt(self, stmt: OpStmt) -> dict:
|
||||
@@ -163,7 +166,7 @@ class MidasAstJsonSerializer(
|
||||
def visit_complex_type(self, type: ComplexType) -> dict:
|
||||
return {
|
||||
"_type": "ComplexType",
|
||||
"properties": self._serialize_list(type.properties),
|
||||
"members": self._serialize_list(type.members),
|
||||
}
|
||||
|
||||
def visit_function_type(self, type: FunctionType) -> dict:
|
||||
@@ -180,3 +183,10 @@ class MidasAstJsonSerializer(
|
||||
"type": arg.type.accept(self),
|
||||
"required": arg.required,
|
||||
}
|
||||
|
||||
def visit_extension_type(self, type: ExtensionType) -> dict:
|
||||
return {
|
||||
"_type": "ExtensionType",
|
||||
"base": type.base.accept(self),
|
||||
"extension": type.extension.accept(self),
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user