Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 31 additions & 1 deletion tests/test_eval_call_with_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,15 @@

import pytest

from typing import Callable, Generic, Literal, Self, TypeVar
from typing import (
Callable,
Generic,
Literal,
Self,
TypeVar,
TypeVarTuple,
Unpack,
)

from typemap.type_eval import eval_call_with_types
from typemap_extensions import (
Expand Down Expand Up @@ -360,3 +368,25 @@ def invoke[U](self, x: U) -> C[U]: ...
GetCallableMember[C[int], Literal["invoke"]], C[int], str
)
assert res == C[str]


def test_eval_call_with_types_var_positional_tvt_01():
def f[*Ts](*args: *Ts) -> tuple[int, *Ts, str]: ...

assert eval_call_with_types(f, int, str) == tuple[int, int, str, str]
assert eval_call_with_types(f) == tuple[int, str]


def test_eval_call_with_types_var_positional_tvt_02():
Ts = TypeVarTuple("Ts")
star = Param[Literal["args"], Unpack[Ts], Literal[ParamKind.VAR_POSITIONAL]]

res = eval_call_with_types(
Callable[Params[star], tuple[int, *Ts, str]], int, str
)
assert res == tuple[int, int, str, str]

res = eval_call_with_types(
Callable[Params[star], Callable[[*Ts], int]], int, str
)
assert res == Callable[[int, str], int]
21 changes: 21 additions & 0 deletions tests/test_type_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
Self,
Tuple,
TypeVar,
TypeVarTuple,
Unpack,
Union,
get_args,
overload,
Expand Down Expand Up @@ -2736,3 +2738,22 @@ def test_raise_error_with_literal_types():
eval_typing(
RaiseError[Literal["Shape mismatch"], Literal[4], Literal[3]]
)


def test_substitute_01():
from typemap.type_eval._apply_generic import substitute

U = TypeVar("U")
Ts = TypeVarTuple("Ts")
args = {T: int, Ts: tuple[int, str]}

assert substitute(list[T], args) == list[int]
assert substitute(T | None, args) == int | None
assert substitute(Callable[[T], U], args) == Callable[[int], U]
assert substitute(Annotated[T, "x"], args) == Annotated[int, "x"]
assert substitute(Annotated[T, {}], args) == Annotated[int, {}]
assert substitute(Unpack[Ts], args) == Unpack[tuple[int, str]]
assert substitute(IsAssignable[T, U], args) == IsAssignable[int, U]
assert substitute(tuple[T, *Ts], args) == tuple[int, int, str]
assert substitute(Callable[[*Ts], T], args) == Callable[[int, str], int]
assert substitute(tuple[U, *Ts], {T: int}) == tuple[U, *Ts]
32 changes: 21 additions & 11 deletions typemap/type_eval/_apply_generic.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,6 @@
import types
import typing

from typing import _GenericAlias as typing_GenericAlias # type: ignore [attr-defined] # noqa: PLC2701


from . import _eval_typing
from . import _typing_inspect
Expand Down Expand Up @@ -79,17 +77,29 @@ def dump(self, *, _level: int = 0):
b.dump(_level=_level + 1)


def _subst_arg(param, args):
is_tvt = isinstance(param, typing.TypeVarTuple)
if param not in args:
return typing.Unpack[param] if is_tvt else param
arg = args[param]
if is_tvt and typing.get_origin(arg) is tuple:
return typing.Unpack[arg]
return arg


def substitute(ty, args):
if ty in args:
return args[ty]
elif isinstance(
ty, (typing_GenericAlias, types.GenericAlias, types.UnionType)
):
return ty.__origin__[*[substitute(t, args) for t in ty.__args__]]
elif isinstance(ty, list):
return [substitute(t, args) for t in ty]
else:
if isinstance(ty, (typing.TypeVar, typing.TypeVarTuple, typing.ParamSpec)):
return args.get(ty, ty)
elif typing.get_origin(ty) is typing.Unpack:
return typing.Unpack[substitute(typing.get_args(ty)[0], args)]

# Lean on typing's own substitution machinery, which knows how to
# handle Callable, Annotated, etc., by wrapping ty in a tuple.
wrapped: Any = tuple[ty]
params = wrapped.__parameters__
if not any(p in args for p in params):
return ty
return wrapped[*[_subst_arg(p, args) for p in params]].__args__[0]


def box(cls: type[Any]) -> Boxed:
Expand Down
Loading