Fallbacks for typing.overload

Would an @overload-related fallback decorator that catches everything not assignable to an existing overload be feasible? For example, consider the following:

from typing import overload

@overload
def spam(ham: int) -> str: ...
@overload
def spam(ham: str) -> int: ...
@overload_fallback
def spam(ham: Any) -> None: ...
def spam(ham: int | str | bytes | bytearray) -> str | int | None: ...

Then spam(b"12") would be inferred as None since (ham: bytes) is not assignable to (ham: int) or (ham: str), but it is assignable to (ham: int | str | bytes | bytearray). Meanwhile, spam([0]) would be a type-checker error since [0] is not int | str | bytes | bytearray.

More precisely, if a call to an overloaded function cannot possibly match any overloads, but it does match the annotation for the implementation, then its return type would be inferred to the return type of the @overload_fallback definition.

The motivating case for this idea is essentially this Stack Overflow question, in which a user wanted to return a float when all four arguments to a function are float and np.ndarray if any of the arguments are np.ndarray. In that case, a correct annotation would require four extra overloads; you can cheat with a type checker directive, but that messes up inference.

(I think this specification is still light on precision since there’s some weirdness around union arguments.)

Shouldn’t this be None then? EDIT: now fixed in the OP

1 Like

I’m against this. I’m going to just deal with the original example, and not the constructed one as the original example contains a stronger reason not to use this as a solution.

Overloads in the original question can be removed entirely and replaced with:


T = TypeVar("T", float, float | np.ndarray, np.ndarray)

def foo(a: T, b: T, c: T, d: T) -> T:
    return a * b + c * d

Any reasonable usage of the function will be correct, and will not have any false negatives. This can be made “more precise” by expanding overloads, but it requires more overloads than the original question suggests, and more precise does not mean more correct; Any usage of such a function that doesn’t know it has only one of these types wouldn’t be able to rely on the result being more precise, so the union showing up in some cases where it could be more precisely known is limited to cases where the function using it is blind to having more precise input knowledge.

Looking at using a fallback rather than overload resolution, we can find a problem:

@overload
def foo(a: float, b: float, c: float, d: float) -> float:  ...
@overload_fallback
def foo(a: Any, b: Any, c: Any, d: Any) -> np.ndarray:  ...
def foo(a: float | np.ndarray, b: float | np.ndarray, c: float | np.ndarray, d: float | np.ndarray) -> float | np.ndarray:
    return a * b + c * d


def failure(a: float | np.ndarray, b: float, c: float, d: float):
    reveal_type(foo(a,b,c,d))  # should be float | np.ndarray, fallback behavior gets np.ndarray

In that case, the result is not necessarily an np.ndarray, and the fallback use would produce the wrong result.

2 Likes

For completeness sake, the constraints can be combined with overloads to increase precision without costing correctness and reducing the overall number of overloads needed to retain correctness, but I can’t imagine a calling function that would benefit from the increased precision.

This is done by removing the union from the constraints, providing the constraint use as an overload, and 1 overload for each parameter where the parameter is known to be distinctly np.ndarray, returning np.ndarray

The resolution complexity for this is unchanged from fully expanding this as overloads, but it involves writing less of it yourself and relying more on overload resolution rules.

T = TypeVar("T", float, np.ndarray)

@overload
def foo(a: T, b: T, c: T, d: T) -> T:  ...
@overload
def foo(a: np.ndarray, b: float | np.ndarray, c: float | np.ndarray, d: float | np.ndarray) -> np.ndarray:  ...
@overload
def foo(a: float | np.ndarray, b: np.ndarray, c: float | np.ndarray, d: float | np.ndarray) -> np.ndarray:  ...
@overload
def foo(a: float | np.ndarray, b: float | np.ndarray, c: np.ndarray, d: float | np.ndarray) -> np.ndarray:  ...
@overload
def foo(a: float | np.ndarray, b: float | np.ndarray, c: float | np.ndarray, d: np.ndarray) -> np.ndarray:  ...
def foo(a: float | np.ndarray, b: float | np.ndarray, c: float | np.ndarray, d: float | np.ndarray) -> float | np.ndarray:
    return a * b + c * d
    ...

pyright example using a provided X and Y class, as this is a generalizable idea

1 Like

IMO this counter-example does not imply the fallback idea is wrong. It just shows the type inferring mechanism must be improved to process fallback overloads correctly.

Without fallback we have to write many combinations of overloads. With fallback the inferring mechanism has to find these combinations by itself. The required work does not disappear. It just has to be done elsewhere.

How do you define “Correctly” in a way that users actually get to write fewer overloads and still get the right result? The idea of a fallback is “when no overload matches, use this”, but the motivating example doesn’t benefit from it, because so many other cases have to be eliminated first for such a fallback to be correct. The work to write those cases as overloads doesn’t go away.

The inference behavior is already doing the heavy lifting in both the simplified and correct version I’d reccomend, and the one that is more precise that I would suggest is unnecessary for most real code; The number of overloads needed if not for leaning on various parts of overload resolution like union expansion rules would be 81