From 43ec29a766dfb80dade75b279094fd4f1325f8df Mon Sep 17 00:00:00 2001 From: "Michael J. Sullivan" Date: Wed, 30 Sep 2026 13:00:07 -0700 Subject: [PATCH] Use `typing`'s built-in substitution stuff support Opus 5.5 thought of this trick. It catches a bunch of edge cases that our recursive `substitute` missed (Callable, Annotated, stuff involving unpacking). --- tests/test_eval_call_with_types.py | 32 ++++++++++++++++++++++++++++- tests/test_type_eval.py | 21 +++++++++++++++++++ typemap/type_eval/_apply_generic.py | 32 +++++++++++++++++++---------- 3 files changed, 73 insertions(+), 12 deletions(-) diff --git a/tests/test_eval_call_with_types.py b/tests/test_eval_call_with_types.py index 55d4c8c1..91e6ea33 100644 --- a/tests/test_eval_call_with_types.py +++ b/tests/test_eval_call_with_types.py @@ -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 ( @@ -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] diff --git a/tests/test_type_eval.py b/tests/test_type_eval.py index c6d8fbd4..6c32413b 100644 --- a/tests/test_type_eval.py +++ b/tests/test_type_eval.py @@ -16,6 +16,8 @@ Self, Tuple, TypeVar, + TypeVarTuple, + Unpack, Union, get_args, overload, @@ -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] diff --git a/typemap/type_eval/_apply_generic.py b/typemap/type_eval/_apply_generic.py index 2b95fa1a..d841c0d4 100644 --- a/typemap/type_eval/_apply_generic.py +++ b/typemap/type_eval/_apply_generic.py @@ -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 @@ -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: