66 lines
2.5 KiB
Python
66 lines
2.5 KiB
Python
Polarity = Literal[-1, 0, 1]
|
|
class VarianceInferrer:
|
|
def __init__(self, manager: VarianceManager) -> None:
|
|
self.manager: VarianceManager = manager
|
|
self.tracker: Tracker = Tracker([])
|
|
|
|
@property
|
|
def types(self) -> TypesRegistry:
|
|
return self.manager.types
|
|
|
|
def infer(self, type: GenericType) -> GenericType:
|
|
self.tracker = Tracker(type.params)
|
|
self.walk(type.body, 1, type.name)
|
|
members: dict[str, Member] = self.types._members.get(type.name, {})
|
|
for name, member in members.items():
|
|
self.walk(member.type, 1, type.name)
|
|
return GenericType(
|
|
name=type.name,
|
|
params=self.tracker.get_updated_vars(),
|
|
body=type.body,
|
|
)
|
|
def walk(
|
|
self,
|
|
type: Type,
|
|
polarity: Polarity,
|
|
base_name: str,
|
|
):
|
|
match type:
|
|
case Function(params=spec):
|
|
all_params: list[Function.Parameter] = spec.pos + spec.mixed + spec.kw
|
|
for param in all_params:
|
|
self.walk(
|
|
param.type,
|
|
-polarity,
|
|
base_name,
|
|
)
|
|
self.walk(type.returns, polarity, base_name)
|
|
case OverloadedFunction(overloads=overloads):
|
|
for overload in overloads:
|
|
self.walk(overload, polarity, base_name)
|
|
case AppliedType(name=name, args=args):
|
|
if self.manager.is_in_queue(name):
|
|
return
|
|
|
|
generic: Type = self.types.get_type(name)
|
|
assert isinstance(generic, GenericType)
|
|
generic = self.manager.infer(name, generic)
|
|
|
|
params: list[TypeVar] = generic.params
|
|
polarities: dict[Variance, Polarity] = {
|
|
Variance.INVARIANT: 0,
|
|
Variance.COVARIANT: 1,
|
|
Variance.CONTRAVARIANT: -1,
|
|
}
|
|
for arg, param in zip(args, params):
|
|
param_polarity: Polarity = polarities[param.variance]
|
|
self.walk(
|
|
arg,
|
|
cast(Polarity, polarity * param_polarity),
|
|
base_name,
|
|
)
|
|
case ConstraintType(type=base):
|
|
self.walk(base, polarity, base_name)
|
|
case TypeVar():
|
|
if type in self.tracker:
|
|
self.tracker.record(type, polarity) |