In Python, if I have a function which takes a callable with N args, can I type hint the signature to indicate that the returned callable also takes N args (though those args are different types)?
As an example, I'm writing a function to wrap Python functions as PySpark UDFs. For Python functions which take single arguments, the wrapper function's signature would look a bit like this:
from typing import Any, Callable
from pyspark.sql.column import Column
def create_udf(py_func: Callable[[Any], Any]) -> Callable[[Column], Column]:
"""Create a PySpark UDF from a Python function."""
...
However, py_func could actually take two arguments, in which case the wrapper function's signature would look like this:
from typing import Any, Callable
from pyspark.sql.column import Column
def create_udf(
py_func: Callable[[Any, Any], Any]
) -> Callable[[Column, Column], Column]:
"""Create a PySpark UDF from a Python function."""
...
My first thought was to use a typing.Protocol and take variable arguments, but that isn't accurate: the number of arguments is fixed, not variable. The only solution I've found so far which works is to use @overload with each supported number of arguments:
from typing import Any, Callable, overload
from pyspark.sql.column import Column
OneArgWrappable = Callable[[Any], Any]
"""A wrappable function taking a single arg."""
TwoArgWrappable = Callable[[Any, Any], Any]
"""A wrappable function taking two args."""
ThreeArgWrappable = Callable[[Any, Any, Any], Any]
"""A wrappable function taking three args."""
FourArgWrappable = Callable[[Any, Any, Any, Any], Any]
"""A wrappable function taking four args."""
OneArgWrapped = Callable[[Column], Column]
"""A wrapped function (Spark UDF) taking a single arg."""
TwoArgWrapped = Callable[[Column, Column], Column]
"""A wrapped function (Spark UDF) taking two args."""
ThreeArgWrapped = Callable[[Column, Column, Column], Column]
"""A wrapped function (Spark UDF) taking three args."""
FourArgWrapped = Callable[[Column, Column, Column, Column], Column]
"""A wrapped function (Spark UDF) taking four args."""
@overload
def create_udf(py_func: OneArgWrappable) -> OneArgWrapped:
pass
@overload
def create_udf(py_func: TwoArgWrappable) -> TwoArgWrapped:
pass
@overload
def create_udf(py_func: ThreeArgWrappable) -> ThreeArgWrapped:
pass
@overload
def create_udf(py_func: FourArgWrappable) -> FourArgWrapped:
pass
def create_udf(py_func: Callable) -> Callable:
"""Create a PySpark UDF from a Python function."""
...
This works, but obviously it's verbose and repetitive. Is there a better solution?