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