Files
TB-Docs/report/code/variance_inferrer.py

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)