Compare commits
5
Commits
official
...
9dd7801d2d
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9dd7801d2d
|
||
|
|
154cb8b314
|
||
|
|
c64ab434b5
|
||
|
|
25e6410546
|
||
|
|
8a22acc17c
|
+4
-155
@@ -6,19 +6,17 @@ from typing import Optional
|
|||||||
import midas.ast.midas as m
|
import midas.ast.midas as m
|
||||||
import midas.ast.python as p
|
import midas.ast.python as p
|
||||||
from midas.ast.location import Location
|
from midas.ast.location import Location
|
||||||
from midas.checker.builtins import BUILTIN_SUBTYPES
|
|
||||||
from midas.checker.diagnostic import Diagnostic, DiagnosticType
|
from midas.checker.diagnostic import Diagnostic, DiagnosticType
|
||||||
from midas.checker.environment import Environment
|
from midas.checker.environment import Environment
|
||||||
from midas.checker.operators import COMPARATOR_METHODS, OPERATOR_METHODS
|
from midas.checker.operators import COMPARATOR_METHODS, OPERATOR_METHODS
|
||||||
from midas.checker.types import (
|
from midas.checker.types import (
|
||||||
AliasType,
|
|
||||||
BaseType,
|
|
||||||
ComplexType,
|
ComplexType,
|
||||||
Function,
|
Function,
|
||||||
Operation,
|
Operation,
|
||||||
Type,
|
Type,
|
||||||
UnitType,
|
UnitType,
|
||||||
UnknownType,
|
UnknownType,
|
||||||
|
unfold_type,
|
||||||
)
|
)
|
||||||
from midas.lexer.midas import MidasLexer
|
from midas.lexer.midas import MidasLexer
|
||||||
from midas.lexer.token import Token
|
from midas.lexer.token import Token
|
||||||
@@ -178,157 +176,8 @@ class Checker(
|
|||||||
stmts: list[m.Stmt] = parser.parse()
|
stmts: list[m.Stmt] = parser.parse()
|
||||||
self.ctx.resolve(stmts)
|
self.ctx.resolve(stmts)
|
||||||
|
|
||||||
def unfold_type(self, type: Type) -> Type:
|
|
||||||
match type:
|
|
||||||
case AliasType(type=ref_type):
|
|
||||||
return self.unfold_type(ref_type)
|
|
||||||
case _:
|
|
||||||
return type
|
|
||||||
|
|
||||||
def is_subtype(self, type1: Type, type2: Type) -> bool:
|
def is_subtype(self, type1: Type, type2: Type) -> bool:
|
||||||
"""Check whether `type1` is a subtype of `type2`
|
return self.ctx.is_subtype(type1, type2)
|
||||||
|
|
||||||
For more details on the rules checked here, see TAPL Chap. 15-16-17
|
|
||||||
|
|
||||||
Args:
|
|
||||||
type1 (Type): the potential subtype
|
|
||||||
type2 (Type): the potential supertype
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
bool: whether `type1` is a subtype of `type2`
|
|
||||||
"""
|
|
||||||
|
|
||||||
if type1 == type2:
|
|
||||||
return True
|
|
||||||
|
|
||||||
match (type1, type2):
|
|
||||||
case (AliasType(type=base1), _):
|
|
||||||
return self.is_subtype(base1, type2)
|
|
||||||
|
|
||||||
case (BaseType(name=name1), BaseType(name=name2)):
|
|
||||||
return name1 in BUILTIN_SUBTYPES.get(name2, set())
|
|
||||||
|
|
||||||
case (ComplexType(properties=props1), ComplexType(properties=props2)):
|
|
||||||
for k, t in props2.items():
|
|
||||||
if k not in props1:
|
|
||||||
return False
|
|
||||||
if not self.is_subtype(props1[k], t):
|
|
||||||
return False
|
|
||||||
return True
|
|
||||||
|
|
||||||
case (Function(returns=return1), Function(returns=return2)):
|
|
||||||
if not self.is_func_subtype(type1, type2):
|
|
||||||
return False
|
|
||||||
if not self.is_subtype(return1, return2):
|
|
||||||
return False
|
|
||||||
return True
|
|
||||||
|
|
||||||
return False
|
|
||||||
|
|
||||||
# TODO: verify the logic in here
|
|
||||||
def is_func_subtype(self, func1: Function, func2: Function) -> bool:
|
|
||||||
"""Check whether a function is a subtype of another
|
|
||||||
|
|
||||||
Args:
|
|
||||||
func1 (Function): the potential function subtype
|
|
||||||
func2 (Function): the potential function supertype
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
bool: whether `func1` is a subtype of `func2`
|
|
||||||
"""
|
|
||||||
if not self.is_subtype(func1.returns, func2.returns):
|
|
||||||
return False
|
|
||||||
|
|
||||||
pos1: list[Function.Argument] = func1.pos_args
|
|
||||||
mixed1: list[Function.Argument] = func1.args
|
|
||||||
kw1: dict[str, Function.Argument] = {a.name: a for a in func1.kw_args}
|
|
||||||
pos2: list[Function.Argument] = func2.pos_args
|
|
||||||
mixed2: list[Function.Argument] = func2.args
|
|
||||||
kw2: dict[str, Function.Argument] = {a.name: a for a in func2.kw_args}
|
|
||||||
|
|
||||||
mixed_by_pos: dict[int, Function.Argument] = {arg.pos: arg for arg in mixed2}
|
|
||||||
mixed_by_name: dict[str, Function.Argument] = {arg.name: arg for arg in mixed2}
|
|
||||||
|
|
||||||
def is_arg_subtype(sub: Function.Argument, sup: Function.Argument) -> bool:
|
|
||||||
if not self.is_subtype(sub.type, sup.type):
|
|
||||||
return False
|
|
||||||
if not sup.required and sub.required:
|
|
||||||
return False
|
|
||||||
return True
|
|
||||||
|
|
||||||
for arg1 in pos1:
|
|
||||||
arg2: Function.Argument
|
|
||||||
if arg1.pos < len(pos2):
|
|
||||||
arg2 = pos2[arg1.pos]
|
|
||||||
elif arg1.pos in mixed_by_pos:
|
|
||||||
arg2 = mixed_by_pos[arg1.pos]
|
|
||||||
elif not arg1.required:
|
|
||||||
continue
|
|
||||||
else:
|
|
||||||
return False
|
|
||||||
if not is_arg_subtype(arg2, arg1):
|
|
||||||
return False
|
|
||||||
|
|
||||||
for name, arg1 in kw1.items():
|
|
||||||
arg2: Function.Argument
|
|
||||||
if name in kw2:
|
|
||||||
arg2 = kw2[name]
|
|
||||||
elif name in mixed_by_name:
|
|
||||||
arg2 = mixed_by_name[name]
|
|
||||||
elif not arg1.required:
|
|
||||||
continue
|
|
||||||
else:
|
|
||||||
return False
|
|
||||||
if not is_arg_subtype(arg2, arg1):
|
|
||||||
return False
|
|
||||||
|
|
||||||
for arg1 in mixed1:
|
|
||||||
pos_arg2: Optional[Function.Argument] = None
|
|
||||||
kw_arg2: Optional[Function.Argument] = None
|
|
||||||
if arg1.name in kw2:
|
|
||||||
kw_arg2 = kw2[arg1.name]
|
|
||||||
elif arg1.name in mixed_by_name:
|
|
||||||
kw_arg2 = mixed_by_name[arg1.name]
|
|
||||||
if arg1.pos < len(pos2):
|
|
||||||
pos_arg2 = pos2[arg1.pos]
|
|
||||||
elif arg1.pos in mixed_by_pos:
|
|
||||||
pos_arg2 = mixed_by_pos[arg1.pos]
|
|
||||||
|
|
||||||
# No match in func2 and arg is required
|
|
||||||
if pos_arg2 is None and kw_arg2 is None and arg1.required:
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Matching keyword argument
|
|
||||||
if kw_arg2 is not None and not is_arg_subtype(kw_arg2, arg1):
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Matching positional argument
|
|
||||||
if pos_arg2 is not None and not is_arg_subtype(pos_arg2, arg1):
|
|
||||||
return False
|
|
||||||
|
|
||||||
mixed_positions: set[int] = {a.pos for a in mixed1}
|
|
||||||
mixed_names: set[str] = {a.name for a in mixed1}
|
|
||||||
for arg2 in pos2:
|
|
||||||
if not arg2.required:
|
|
||||||
continue
|
|
||||||
if arg2.pos >= len(pos1) and arg2.pos not in mixed_positions:
|
|
||||||
return False
|
|
||||||
|
|
||||||
for name, arg2 in kw2.items():
|
|
||||||
if not arg2.required:
|
|
||||||
continue
|
|
||||||
if name not in kw1 and name not in mixed_names:
|
|
||||||
return False
|
|
||||||
|
|
||||||
for arg2 in mixed2:
|
|
||||||
if arg2.required:
|
|
||||||
continue
|
|
||||||
pos_match: bool = arg2.pos < len(pos1) or arg2.pos in mixed_positions
|
|
||||||
kw_match: bool = arg2.name in kw1 or arg2.name in mixed_names
|
|
||||||
if not pos_match or not kw_match:
|
|
||||||
return False
|
|
||||||
|
|
||||||
return True
|
|
||||||
|
|
||||||
def visit_expression_stmt(self, stmt: p.ExpressionStmt) -> None:
|
def visit_expression_stmt(self, stmt: p.ExpressionStmt) -> None:
|
||||||
self.type_of(stmt.expr)
|
self.type_of(stmt.expr)
|
||||||
@@ -470,7 +319,7 @@ class Checker(
|
|||||||
|
|
||||||
def _assign_attr(self, location: Location, target: p.GetExpr, value_type: Type):
|
def _assign_attr(self, location: Location, target: p.GetExpr, value_type: Type):
|
||||||
object: Type = self.type_of(target.object)
|
object: Type = self.type_of(target.object)
|
||||||
base_object: Type = self.unfold_type(object)
|
base_object: Type = unfold_type(object)
|
||||||
match base_object:
|
match base_object:
|
||||||
case ComplexType(properties=properties):
|
case ComplexType(properties=properties):
|
||||||
if target.name not in properties:
|
if target.name not in properties:
|
||||||
@@ -611,7 +460,7 @@ class Checker(
|
|||||||
|
|
||||||
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 = self.unfold_type(object)
|
base_object: Type = unfold_type(object)
|
||||||
match base_object:
|
match base_object:
|
||||||
case ComplexType(properties=properties):
|
case ComplexType(properties=properties):
|
||||||
if expr.name not in properties:
|
if expr.name not in properties:
|
||||||
|
|||||||
+81
-1
@@ -1,6 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, kw_only=True)
|
@dataclass(frozen=True, kw_only=True)
|
||||||
@@ -57,4 +58,83 @@ class Operation:
|
|||||||
right: Type
|
right: Type
|
||||||
|
|
||||||
|
|
||||||
Type = BaseType | AliasType | UnknownType | UnitType | Function | ComplexType
|
@dataclass(frozen=True, kw_only=True)
|
||||||
|
class TypeVar:
|
||||||
|
name: str
|
||||||
|
bound: Optional[Type]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, kw_only=True)
|
||||||
|
class GenericType:
|
||||||
|
params: list[TypeVar]
|
||||||
|
body: Type
|
||||||
|
|
||||||
|
|
||||||
|
def substitute_typevars(type: Type, substitutions: dict[str, Type]) -> Type:
|
||||||
|
def sub_argument(arg: Function.Argument):
|
||||||
|
return Function.Argument(
|
||||||
|
pos=arg.pos,
|
||||||
|
name=arg.name,
|
||||||
|
type=substitute_typevars(arg.type, substitutions),
|
||||||
|
required=arg.required,
|
||||||
|
)
|
||||||
|
|
||||||
|
match type:
|
||||||
|
case BaseType(name=name) if name in substitutions:
|
||||||
|
return substitutions[name]
|
||||||
|
|
||||||
|
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] = {
|
||||||
|
name: substitute_typevars(prop, substitutions)
|
||||||
|
for name, prop in properties.items()
|
||||||
|
}
|
||||||
|
return ComplexType(properties=properties2)
|
||||||
|
|
||||||
|
case TypeVar(name=name):
|
||||||
|
if name in substitutions:
|
||||||
|
return substitutions[name]
|
||||||
|
raise ValueError(f"Missing TypeVar substitution for {name}")
|
||||||
|
|
||||||
|
case UnknownType() | UnitType():
|
||||||
|
return type
|
||||||
|
|
||||||
|
case _:
|
||||||
|
raise NotImplementedError(f"Unsupported type {type}")
|
||||||
|
|
||||||
|
|
||||||
|
def unfold_type(type: Type) -> Type:
|
||||||
|
match type:
|
||||||
|
case AliasType(type=ref_type):
|
||||||
|
return unfold_type(ref_type)
|
||||||
|
case _:
|
||||||
|
return type
|
||||||
|
|
||||||
|
|
||||||
|
Type = (
|
||||||
|
BaseType
|
||||||
|
| AliasType
|
||||||
|
| UnknownType
|
||||||
|
| UnitType
|
||||||
|
| Function
|
||||||
|
| ComplexType
|
||||||
|
| TypeVar
|
||||||
|
| GenericType
|
||||||
|
)
|
||||||
|
|||||||
+195
-8
@@ -1,12 +1,18 @@
|
|||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
import midas.ast.midas as m
|
import midas.ast.midas as m
|
||||||
|
from midas.checker.builtins import BUILTIN_SUBTYPES
|
||||||
from midas.checker.types import (
|
from midas.checker.types import (
|
||||||
AliasType,
|
AliasType,
|
||||||
|
BaseType,
|
||||||
ComplexType,
|
ComplexType,
|
||||||
|
Function,
|
||||||
|
GenericType,
|
||||||
Operation,
|
Operation,
|
||||||
Type,
|
Type,
|
||||||
|
TypeVar,
|
||||||
UnknownType,
|
UnknownType,
|
||||||
|
substitute_typevars,
|
||||||
)
|
)
|
||||||
from midas.resolver.builtin import define_builtins
|
from midas.resolver.builtin import define_builtins
|
||||||
|
|
||||||
@@ -18,6 +24,8 @@ class MidasResolver(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[T
|
|||||||
self._types: dict[str, Type] = {}
|
self._types: dict[str, Type] = {}
|
||||||
self._operations: dict[Operation.CallSignature, Type] = {}
|
self._operations: dict[Operation.CallSignature, Type] = {}
|
||||||
|
|
||||||
|
self._local_variables: dict[str, TypeVar] = {}
|
||||||
|
|
||||||
define_builtins(self)
|
define_builtins(self)
|
||||||
|
|
||||||
def get_type(self, name: str) -> Type:
|
def get_type(self, name: str) -> Type:
|
||||||
@@ -32,10 +40,11 @@ class MidasResolver(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[T
|
|||||||
Returns:
|
Returns:
|
||||||
Type: the type
|
Type: the type
|
||||||
"""
|
"""
|
||||||
type: Optional[Type] = self._types.get(name)
|
if name in self._local_variables:
|
||||||
if type is None:
|
return self._local_variables[name]
|
||||||
raise NameError(f"Undefined type {name}")
|
if name in self._types:
|
||||||
return type
|
return self._types[name]
|
||||||
|
raise NameError(f"Undefined type {name}")
|
||||||
|
|
||||||
def get_operation_result(
|
def get_operation_result(
|
||||||
self, left: Type, operator: str, right: Type
|
self, left: Type, operator: str, right: Type
|
||||||
@@ -121,12 +130,21 @@ class MidasResolver(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[T
|
|||||||
stmt.accept(self)
|
stmt.accept(self)
|
||||||
|
|
||||||
def visit_type_stmt(self, stmt: m.TypeStmt) -> None:
|
def visit_type_stmt(self, stmt: m.TypeStmt) -> None:
|
||||||
type: Type = stmt.type.accept(self)
|
params: list[TypeVar] = []
|
||||||
for param in stmt.params:
|
for param in stmt.params:
|
||||||
|
name: str = param.name.lexeme
|
||||||
|
bound: Optional[Type] = None
|
||||||
if param.bound is not None:
|
if param.bound is not None:
|
||||||
param.bound.accept(self)
|
bound = param.bound.accept(self)
|
||||||
|
var = TypeVar(name=name, bound=bound)
|
||||||
|
self._local_variables[name] = var
|
||||||
|
params.append(var)
|
||||||
|
type: Type = stmt.type.accept(self)
|
||||||
|
if len(params) != 0:
|
||||||
|
type = GenericType(params=params, body=type)
|
||||||
name: str = stmt.name.lexeme
|
name: str = stmt.name.lexeme
|
||||||
self.define_type(name, AliasType(name=name, type=type))
|
self.define_type(name, AliasType(name=name, type=type))
|
||||||
|
self._local_variables.clear()
|
||||||
|
|
||||||
def visit_property_stmt(self, stmt: m.PropertyStmt) -> None: ...
|
def visit_property_stmt(self, stmt: m.PropertyStmt) -> None: ...
|
||||||
|
|
||||||
@@ -169,8 +187,36 @@ class MidasResolver(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[T
|
|||||||
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)
|
||||||
params: list[Type] = [param.accept(self) for param in type.params]
|
params: list[Type] = [param.accept(self) for param in type.params]
|
||||||
# TODO
|
return self.apply_generic(type_, params)
|
||||||
return UnknownType()
|
|
||||||
|
def apply_generic(self, type: Type, params: list[Type]) -> Type:
|
||||||
|
match type:
|
||||||
|
case AliasType(name=name, type=base):
|
||||||
|
return AliasType(name=name, type=self.apply_generic(base, params))
|
||||||
|
|
||||||
|
case GenericType(params=type_vars, body=body):
|
||||||
|
n_params: int = len(params)
|
||||||
|
n_type_vars: int = len(type_vars)
|
||||||
|
if n_params < n_type_vars:
|
||||||
|
raise ValueError(
|
||||||
|
f"Missing type parameters, expected {n_type_vars} but only {n_params} provided"
|
||||||
|
)
|
||||||
|
if n_params > n_type_vars:
|
||||||
|
raise ValueError(
|
||||||
|
f"Too many type parameters, expected {n_type_vars} but {n_params} provided"
|
||||||
|
)
|
||||||
|
substitutions: dict[str, Type] = {}
|
||||||
|
for param, type_var in zip(params, type_vars):
|
||||||
|
if type_var.bound is not None and not self.is_subtype(
|
||||||
|
param, type_var.bound
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"Type parameter {param} is not a subtype of {type_var.bound}"
|
||||||
|
)
|
||||||
|
substitutions[type_var.name] = param
|
||||||
|
return substitute_typevars(body, substitutions)
|
||||||
|
case _:
|
||||||
|
raise ValueError(f"{type} is not a generic type")
|
||||||
|
|
||||||
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)
|
||||||
@@ -184,3 +230,144 @@ class MidasResolver(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[T
|
|||||||
prop.name.lexeme: prop.type.accept(self) for prop in type.properties
|
prop.name.lexeme: prop.type.accept(self) for prop in type.properties
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def is_subtype(self, type1: Type, type2: Type) -> bool:
|
||||||
|
"""Check whether `type1` is a subtype of `type2`
|
||||||
|
|
||||||
|
For more details on the rules checked here, see TAPL Chap. 15-16-17
|
||||||
|
|
||||||
|
Args:
|
||||||
|
type1 (Type): the potential subtype
|
||||||
|
type2 (Type): the potential supertype
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: whether `type1` is a subtype of `type2`
|
||||||
|
"""
|
||||||
|
|
||||||
|
if type1 == type2:
|
||||||
|
return True
|
||||||
|
|
||||||
|
match (type1, type2):
|
||||||
|
case (AliasType(type=base1), _):
|
||||||
|
return self.is_subtype(base1, type2)
|
||||||
|
|
||||||
|
case (BaseType(name=name1), BaseType(name=name2)):
|
||||||
|
return name1 in BUILTIN_SUBTYPES.get(name2, set())
|
||||||
|
|
||||||
|
case (ComplexType(properties=props1), ComplexType(properties=props2)):
|
||||||
|
for k, t in props2.items():
|
||||||
|
if k not in props1:
|
||||||
|
return False
|
||||||
|
if not self.is_subtype(props1[k], t):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
case (Function(), Function()):
|
||||||
|
return self.is_func_subtype(type1, type2)
|
||||||
|
|
||||||
|
return False
|
||||||
|
|
||||||
|
# TODO: verify the logic in here
|
||||||
|
def is_func_subtype(self, func1: Function, func2: Function) -> bool:
|
||||||
|
"""Check whether a function is a subtype of another
|
||||||
|
|
||||||
|
Args:
|
||||||
|
func1 (Function): the potential function subtype
|
||||||
|
func2 (Function): the potential function supertype
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: whether `func1` is a subtype of `func2`
|
||||||
|
"""
|
||||||
|
if not self.is_subtype(func1.returns, func2.returns):
|
||||||
|
return False
|
||||||
|
|
||||||
|
pos1: list[Function.Argument] = func1.pos_args
|
||||||
|
mixed1: list[Function.Argument] = func1.args
|
||||||
|
kw1: dict[str, Function.Argument] = {a.name: a for a in func1.kw_args}
|
||||||
|
pos2: list[Function.Argument] = func2.pos_args
|
||||||
|
mixed2: list[Function.Argument] = func2.args
|
||||||
|
kw2: dict[str, Function.Argument] = {a.name: a for a in func2.kw_args}
|
||||||
|
|
||||||
|
mixed_by_pos: dict[int, Function.Argument] = {arg.pos: arg for arg in mixed2}
|
||||||
|
mixed_by_name: dict[str, Function.Argument] = {arg.name: arg for arg in mixed2}
|
||||||
|
|
||||||
|
def is_arg_subtype(sub: Function.Argument, sup: Function.Argument) -> bool:
|
||||||
|
if not self.is_subtype(sub.type, sup.type):
|
||||||
|
return False
|
||||||
|
if not sup.required and sub.required:
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
for arg1 in pos1:
|
||||||
|
arg2: Function.Argument
|
||||||
|
if arg1.pos < len(pos2):
|
||||||
|
arg2 = pos2[arg1.pos]
|
||||||
|
elif arg1.pos in mixed_by_pos:
|
||||||
|
arg2 = mixed_by_pos[arg1.pos]
|
||||||
|
elif not arg1.required:
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
return False
|
||||||
|
if not is_arg_subtype(arg2, arg1):
|
||||||
|
return False
|
||||||
|
|
||||||
|
for name, arg1 in kw1.items():
|
||||||
|
arg2: Function.Argument
|
||||||
|
if name in kw2:
|
||||||
|
arg2 = kw2[name]
|
||||||
|
elif name in mixed_by_name:
|
||||||
|
arg2 = mixed_by_name[name]
|
||||||
|
elif not arg1.required:
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
return False
|
||||||
|
if not is_arg_subtype(arg2, arg1):
|
||||||
|
return False
|
||||||
|
|
||||||
|
for arg1 in mixed1:
|
||||||
|
pos_arg2: Optional[Function.Argument] = None
|
||||||
|
kw_arg2: Optional[Function.Argument] = None
|
||||||
|
if arg1.name in kw2:
|
||||||
|
kw_arg2 = kw2[arg1.name]
|
||||||
|
elif arg1.name in mixed_by_name:
|
||||||
|
kw_arg2 = mixed_by_name[arg1.name]
|
||||||
|
if arg1.pos < len(pos2):
|
||||||
|
pos_arg2 = pos2[arg1.pos]
|
||||||
|
elif arg1.pos in mixed_by_pos:
|
||||||
|
pos_arg2 = mixed_by_pos[arg1.pos]
|
||||||
|
|
||||||
|
# No match in func2 and arg is required
|
||||||
|
if pos_arg2 is None and kw_arg2 is None and arg1.required:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Matching keyword argument
|
||||||
|
if kw_arg2 is not None and not is_arg_subtype(kw_arg2, arg1):
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Matching positional argument
|
||||||
|
if pos_arg2 is not None and not is_arg_subtype(pos_arg2, arg1):
|
||||||
|
return False
|
||||||
|
|
||||||
|
mixed_positions: set[int] = {a.pos for a in mixed1}
|
||||||
|
mixed_names: set[str] = {a.name for a in mixed1}
|
||||||
|
for arg2 in pos2:
|
||||||
|
if not arg2.required:
|
||||||
|
continue
|
||||||
|
if arg2.pos >= len(pos1) and arg2.pos not in mixed_positions:
|
||||||
|
return False
|
||||||
|
|
||||||
|
for name, arg2 in kw2.items():
|
||||||
|
if not arg2.required:
|
||||||
|
continue
|
||||||
|
if name not in kw1 and name not in mixed_names:
|
||||||
|
return False
|
||||||
|
|
||||||
|
for arg2 in mixed2:
|
||||||
|
if arg2.required:
|
||||||
|
continue
|
||||||
|
pos_match: bool = arg2.pos < len(pos1) or arg2.pos in mixed_positions
|
||||||
|
kw_match: bool = arg2.name in kw1 or arg2.name in mixed_names
|
||||||
|
if not pos_match or not kw_match:
|
||||||
|
return False
|
||||||
|
|
||||||
|
return True
|
||||||
|
|||||||
Reference in New Issue
Block a user