28 lines
940 B
Python
28 lines
940 B
Python
class Tracker:
|
|
def __init__(self, vars: list[TypeVar]) -> None:
|
|
self.vars: list[TypeVar] = vars
|
|
self.refs: dict[str, set[Polarity]] = {var.name: set() for var in self.vars}
|
|
|
|
def record(self, var: TypeVar, polarity: Polarity):
|
|
self.refs[var.name].add(polarity)
|
|
|
|
def get_updated_vars(self) -> list[TypeVar]:
|
|
return [
|
|
TypeVar(
|
|
name=var.name, bound=var.bound, variance=self.get_variance(var.name)
|
|
)
|
|
for var in self.vars
|
|
]
|
|
|
|
def get_variance(self, name: str) -> Variance:
|
|
refs: set[Polarity] = self.refs[name]
|
|
if refs == {-1}:
|
|
return Variance.CONTRAVARIANT
|
|
if refs == {1}:
|
|
return Variance.COVARIANT
|
|
return Variance.INVARIANT
|
|
|
|
def __contains__(self, item: TypeVar | str):
|
|
if isinstance(item, TypeVar):
|
|
return item.name in self
|
|
return item in self.refs |