15 Commits
12 changed files with 438 additions and 189 deletions
+16 -4
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
+82
View File
@@ -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
View File
@@ -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
+11 -7
View File
@@ -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
View File
@@ -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:
+6
View File
@@ -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
View File
@@ -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)
+16 -6
View File
@@ -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),
}