From 92dedae132e0142fc33a6d38d33473af2186b0e5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=9A=D0=B8=D1=80=D0=B8=D0=BB=D0=BB=20=D0=92=D0=B5=D1=82?= =?UTF-8?q?=D1=80=D0=BE=D0=B2?= Date: Sun, 9 Aug 2026 11:56:59 +0300 Subject: [PATCH 1/3] Fix lambda contextual typing with keyword-only callables --- .../src/analyzer/typeEvaluator.ts | 8 +++- .../src/tests/samples/lambda4.py | 37 ++++++++++++++++++- .../src/tests/typeEvaluator1.test.ts | 2 +- 3 files changed, 44 insertions(+), 3 deletions(-) diff --git a/packages/pyright-internal/src/analyzer/typeEvaluator.ts b/packages/pyright-internal/src/analyzer/typeEvaluator.ts index 7f759ee2af06..2ea0946746ff 100644 --- a/packages/pyright-internal/src/analyzer/typeEvaluator.ts +++ b/packages/pyright-internal/src/analyzer/typeEvaluator.ts @@ -15264,6 +15264,7 @@ export function createTypeEvaluator( // more sophisticated in the future, but it becomes very complex to handle // all of the permutations. let sawParamMismatch = false; + let sawLambdaArgsParam = false; node.d.params.forEach((param, index) => { let paramType: Type | undefined; @@ -15277,7 +15278,8 @@ export function createTypeEvaluator( // from the expected parameter. if ( expectedParam.param.category === param.d.category && - !param.d.name === !expectedParam.param.name + !param.d.name === !expectedParam.param.name && + (expectedParam.kind !== ParamKind.Keyword || sawLambdaArgsParam) ) { paramType = expectedParam.type; } else { @@ -15348,6 +15350,10 @@ export function createTypeEvaluator( ); FunctionType.addParam(functionType, functionParam); + + if (param.d.category === ParamCategory.ArgsList) { + sawLambdaArgsParam = true; + } }); if (paramsArePositionOnly && functionType.shared.parameters.length > 0) { diff --git a/packages/pyright-internal/src/tests/samples/lambda4.py b/packages/pyright-internal/src/tests/samples/lambda4.py index 95269e677f1c..88d2815c1980 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 U1 = Callable[[int, str], bool] | Callable[[str], bool] @@ -76,3 +76,38 @@ 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 diff --git a/packages/pyright-internal/src/tests/typeEvaluator1.test.ts b/packages/pyright-internal/src/tests/typeEvaluator1.test.ts index 73306b3b11a4..3e6ec468604e 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, 3); }); test('Lambda5', () => { From 32b885da73ecb3d76235dddb2fa72f853cfc2c9b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=9A=D0=B8=D1=80=D0=B8=D0=BB=D0=BB=20=D0=92=D0=B5=D1=82?= =?UTF-8?q?=D1=80=D0=BE=D0=B2?= Date: Mon, 10 Aug 2026 07:36:46 +0300 Subject: [PATCH 2/3] Preserve matching keyword lambda context --- .../src/analyzer/typeEvaluator.ts | 15 ++++++++++++++- .../src/tests/samples/lambda4.py | 19 ++++++++++++++++++- .../src/tests/typeEvaluator1.test.ts | 2 +- 3 files changed, 33 insertions(+), 3 deletions(-) diff --git a/packages/pyright-internal/src/analyzer/typeEvaluator.ts b/packages/pyright-internal/src/analyzer/typeEvaluator.ts index 2ea0946746ff..ecdb189dd630 100644 --- a/packages/pyright-internal/src/analyzer/typeEvaluator.ts +++ b/packages/pyright-internal/src/analyzer/typeEvaluator.ts @@ -15265,6 +15265,9 @@ export function createTypeEvaluator( // all of the permutations. let sawParamMismatch = false; let sawLambdaArgsParam = false; + 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; @@ -15272,6 +15275,16 @@ export function createTypeEvaluator( if (expectedParamDetails && !sawParamMismatch) { if (index < expectedParamDetails.params.length) { const expectedParam = expectedParamDetails.params[index]; + const isPositionOnlyParam = + (positionOnlySeparatorIndex >= 0 && index < positionOnlySeparatorIndex) || + (positionOnlySeparatorIndex < 0 && + paramsArePositionOnly && + !!param.d.name && + isPrivateName(param.d.name.d.value)); + const isCompatibleKeywordParam = + expectedParam.kind !== ParamKind.Keyword || + sawLambdaArgsParam || + (!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 @@ -15279,7 +15292,7 @@ export function createTypeEvaluator( if ( expectedParam.param.category === param.d.category && !param.d.name === !expectedParam.param.name && - (expectedParam.kind !== ParamKind.Keyword || sawLambdaArgsParam) + isCompatibleKeywordParam ) { paramType = expectedParam.type; } else { diff --git a/packages/pyright-internal/src/tests/samples/lambda4.py b/packages/pyright-internal/src/tests/samples/lambda4.py index 88d2815c1980..8a63f21bbe0f 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, Generic, Protocol, Self, TypeVar +from typing import Callable, Generic, Protocol, Self, TypeVar, assert_type U1 = Callable[[int, str], bool] | Callable[[str], bool] @@ -111,3 +111,20 @@ def __call__(self, value: int) -> Self: ... # 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) + +# 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 3e6ec468604e..d055123e8c3f 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, 3); + TestUtils.validateResults(analysisResults, 4); }); test('Lambda5', () => { From e016c036592cf972bac48f3f8e2d2df84c816f0e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=9A=D0=B8=D1=80=D0=B8=D0=BB=D0=BB=20=D0=92=D0=B5=D1=82?= =?UTF-8?q?=D1=80=D0=BE=D0=B2?= Date: Mon, 24 Aug 2026 20:53:00 +0300 Subject: [PATCH 3/3] Fix lambda keyword-only parameter matching --- .../src/analyzer/typeEvaluator.ts | 22 ++++++++++++------- .../src/tests/samples/lambda4.py | 8 +++++++ .../src/tests/typeEvaluator1.test.ts | 2 +- 3 files changed, 23 insertions(+), 9 deletions(-) diff --git a/packages/pyright-internal/src/analyzer/typeEvaluator.ts b/packages/pyright-internal/src/analyzer/typeEvaluator.ts index ecdb189dd630..230840ac6b9d 100644 --- a/packages/pyright-internal/src/analyzer/typeEvaluator.ts +++ b/packages/pyright-internal/src/analyzer/typeEvaluator.ts @@ -15264,7 +15264,7 @@ export function createTypeEvaluator( // more sophisticated in the future, but it becomes very complex to handle // all of the permutations. let sawParamMismatch = false; - let sawLambdaArgsParam = false; + let expectedParamIndex = 0; const positionOnlySeparatorIndex = node.d.params.findIndex( (param) => param.d.category === ParamCategory.Simple && !param.d.name ); @@ -15273,8 +15273,17 @@ export function createTypeEvaluator( 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 && @@ -15283,7 +15292,6 @@ export function createTypeEvaluator( isPrivateName(param.d.name.d.value)); const isCompatibleKeywordParam = expectedParam.kind !== ParamKind.Keyword || - sawLambdaArgsParam || (!isPositionOnlyParam && param.d.name?.d.value === expectedParam.param.name); // If the parameter category matches and both of the parameters are @@ -15298,6 +15306,8 @@ export function createTypeEvaluator( } 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 @@ -15363,10 +15373,6 @@ export function createTypeEvaluator( ); FunctionType.addParam(functionType, functionParam); - - if (param.d.category === ParamCategory.ArgsList) { - sawLambdaArgsParam = true; - } }); if (paramsArePositionOnly && functionType.shared.parameters.length > 0) { diff --git a/packages/pyright-internal/src/tests/samples/lambda4.py b/packages/pyright-internal/src/tests/samples/lambda4.py index 8a63f21bbe0f..53f1651f543b 100644 --- a/packages/pyright-internal/src/tests/samples/lambda4.py +++ b/packages/pyright-internal/src/tests/samples/lambda4.py @@ -126,5 +126,13 @@ 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 d055123e8c3f..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, 4); + TestUtils.validateResults(analysisResults, 5); }); test('Lambda5', () => {