diff --git a/packages/pyright-internal/src/analyzer/typeEvaluator.ts b/packages/pyright-internal/src/analyzer/typeEvaluator.ts index 7f759ee2af06..230840ac6b9d 100644 --- a/packages/pyright-internal/src/analyzer/typeEvaluator.ts +++ b/packages/pyright-internal/src/analyzer/typeEvaluator.ts @@ -15264,25 +15264,50 @@ export function createTypeEvaluator( // more sophisticated in the future, but it becomes very complex to handle // all of the permutations. let sawParamMismatch = false; + let expectedParamIndex = 0; + const positionOnlySeparatorIndex = node.d.params.findIndex( + (param) => param.d.category === ParamCategory.Simple && !param.d.name + ); node.d.params.forEach((param, index) => { let paramType: Type | undefined; if (expectedParamDetails && !sawParamMismatch) { - if (index < expectedParamDetails.params.length) { - const expectedParam = expectedParamDetails.params[index]; + const isKeywordOnlySeparator = param.d.category === ParamCategory.ArgsList && !param.d.name; + + if (isKeywordOnlySeparator) { + if ( + expectedParamIndex < expectedParamDetails.params.length && + expectedParamDetails.params[expectedParamIndex].kind !== ParamKind.Keyword + ) { + sawParamMismatch = true; + } + } else if (expectedParamIndex < expectedParamDetails.params.length) { + const expectedParam = expectedParamDetails.params[expectedParamIndex]; + const isPositionOnlyParam = + (positionOnlySeparatorIndex >= 0 && index < positionOnlySeparatorIndex) || + (positionOnlySeparatorIndex < 0 && + paramsArePositionOnly && + !!param.d.name && + isPrivateName(param.d.name.d.value)); + const isCompatibleKeywordParam = + expectedParam.kind !== ParamKind.Keyword || + (!isPositionOnlyParam && param.d.name?.d.value === expectedParam.param.name); // If the parameter category matches and both of the parameters are // either separators (/ or *) or not separators, copy the type // from the expected parameter. if ( expectedParam.param.category === param.d.category && - !param.d.name === !expectedParam.param.name + !param.d.name === !expectedParam.param.name && + isCompatibleKeywordParam ) { paramType = expectedParam.type; } else { sawParamMismatch = true; } + + expectedParamIndex++; } else if (param.d.defaultValue) { // If the lambda param has a default value but there is no associated // parameter in the expected type, assume that the default value is diff --git a/packages/pyright-internal/src/tests/samples/lambda4.py b/packages/pyright-internal/src/tests/samples/lambda4.py index 95269e677f1c..53f1651f543b 100644 --- a/packages/pyright-internal/src/tests/samples/lambda4.py +++ b/packages/pyright-internal/src/tests/samples/lambda4.py @@ -1,7 +1,7 @@ # This sample tests the case where a lambda is assigned to # a union type that contains multiple callables. -from typing import Callable, Protocol, TypeVar +from typing import Callable, Generic, Protocol, Self, TypeVar, assert_type U1 = Callable[[int, str], bool] | Callable[[str], bool] @@ -76,3 +76,63 @@ def accepts_u2(cb: U2) -> U2: def accepts_u3(u: U3): # This should generate an error. u(lambda v: v.lower()) + + +class KeywordOnlyCallable: + def __call__(self, *, kwarg: int) -> Self: ... + + +keyword_only_union: Callable[[KeywordOnlyCallable], KeywordOnlyCallable] | KeywordOnlyCallable = lambda x: x + + +class GenericKeywordOnlyCallable(Generic[T]): + def __call__(self, *, kwarg: T) -> Self: ... + + +generic_keyword_only_union: ( + Callable[[GenericKeywordOnlyCallable[int]], GenericKeywordOnlyCallable[int]] | GenericKeywordOnlyCallable[int] +) = lambda x: x + + +class KeywordOnlyCallback(Protocol): + def __call__(self, *, value: int) -> Self: ... + + +protocol_keyword_only_union: Callable[[KeywordOnlyCallback], KeywordOnlyCallback] | KeywordOnlyCallback = lambda x: x + +ordinary_callable_union: Callable[[int], int] | Callable[[str], str] = lambda x: x + + +class PositionalCallable: + def __call__(self, value: int) -> Self: ... + + +positional_callable_union: Callable[[PositionalCallable], PositionalCallable] | PositionalCallable = lambda x: x + +# This should generate an error. +keyword_only_callback: KeywordOnlyCallback = lambda x: x + + +class KeywordOnlyIntCallback(Protocol): + def __call__(self, *, value: int) -> int: ... + + +same_name_keyword_only_callback: KeywordOnlyIntCallback = lambda value: assert_type(value, int) + + +class VariadicKeywordOnlyCallback(Protocol): + def __call__(self, *args: object, value: int) -> int: ... + + +variadic_keyword_only_callback: VariadicKeywordOnlyCallback = lambda *args, value: assert_type(value, int) + +# The bare `*` separator should not consume the contextual parameter index. +bare_keyword_only_callback: KeywordOnlyIntCallback = lambda *, value: assert_type(value, int) + +# This should generate an error because the keyword-only parameter name differs. +variadic_keyword_only_callback_different_name: VariadicKeywordOnlyCallback = lambda *args, other: reveal_type( + other, expected_text="Unknown" +) + +# This should generate an error. +position_only_keyword_callback: KeywordOnlyIntCallback = lambda value, /: value diff --git a/packages/pyright-internal/src/tests/typeEvaluator1.test.ts b/packages/pyright-internal/src/tests/typeEvaluator1.test.ts index 73306b3b11a4..e526d0caaa9b 100644 --- a/packages/pyright-internal/src/tests/typeEvaluator1.test.ts +++ b/packages/pyright-internal/src/tests/typeEvaluator1.test.ts @@ -726,7 +726,7 @@ test('Lambda3', () => { test('Lambda4', () => { const analysisResults = TestUtils.typeAnalyzeSampleFiles(['lambda4.py']); - TestUtils.validateResults(analysisResults, 2); + TestUtils.validateResults(analysisResults, 5); }); test('Lambda5', () => {