From 348e7b8351a923d10ae69d078da48a85c2d3cfad Mon Sep 17 00:00:00 2001 From: Henry Su Date: Sun, 9 Aug 2026 23:29:12 -0500 Subject: [PATCH 1/3] Fix TypedDict get, pop, and setdefault type evaluation with union keys --- .../src/analyzer/typeEvaluator.ts | 17 ++ .../src/analyzer/typedDicts.ts | 191 ++++++++++++++++++ .../src/tests/samples/typedDict28.py | 45 +++++ .../src/tests/typeEvaluator7.test.ts | 6 + 4 files changed, 259 insertions(+) create mode 100644 packages/pyright-internal/src/tests/samples/typedDict28.py diff --git a/packages/pyright-internal/src/analyzer/typeEvaluator.ts b/packages/pyright-internal/src/analyzer/typeEvaluator.ts index 7f759ee2af06..6e252f57ee38 100644 --- a/packages/pyright-internal/src/analyzer/typeEvaluator.ts +++ b/packages/pyright-internal/src/analyzer/typeEvaluator.ts @@ -177,10 +177,12 @@ import { getLastTypedDeclarationForSymbol, isEffectivelyClassVar } from './symbo import { assignTupleTypeArgs, expandTuple, getSlicedTupleType, getTypeOfTuple, makeTupleObject } from './tuples'; import { SpeculativeModeOptions, SpeculativeTypeTracker } from './typeCacheUtils'; import { + applyTypedDictMethodTransform, assignToTypedDict, assignTypedDictToTypedDict, createTypedDictType, createTypedDictTypeInlined, + getTypedDictClassFromMethod, getTypedDictDictEquivalent, getTypedDictMappingEquivalent, getTypedDictMembersForClass, @@ -10544,6 +10546,21 @@ export function createTypeEvaluator( return { returnType: evaluateCastCall(argList, errorNode) }; } + const tdMethodInfo = getTypedDictClassFromMethod(expandedCallType); + if (tdMethodInfo) { + const tdResult = applyTypedDictMethodTransform( + evaluatorInterface, + errorNode, + argList, + tdMethodInfo.classType, + tdMethodInfo.methodName, + tdMethodInfo.isBound + ); + if (tdResult) { + return tdResult; + } + } + const callResult = validateOverloadedArgTypes( errorNode, argList, diff --git a/packages/pyright-internal/src/analyzer/typedDicts.ts b/packages/pyright-internal/src/analyzer/typedDicts.ts index be0cf4cbdb77..ce625f95c112 100644 --- a/packages/pyright-internal/src/analyzer/typedDicts.ts +++ b/packages/pyright-internal/src/analyzer/typedDicts.ts @@ -34,6 +34,7 @@ import { getLastTypedDeclarationForSymbol } from './symbolUtils'; import { Arg, AssignTypeFlags, + CallResult, EvaluatorUsage, TypeEvaluator, TypeResult, @@ -53,6 +54,8 @@ import { isClassInstance, isInstantiableClass, isNever, + isOverloaded, + isUnion, maxTypeRecursionCount, NeverType, OverloadedType, @@ -1611,6 +1614,194 @@ export function getTypeOfIndexedTypedDict( return { type: resultingType, isIncomplete: !!indexTypeResult.isIncomplete }; } +export function getTypedDictClassFromMethod( + type: FunctionType | OverloadedType +): { classType: ClassType; methodName: string; isBound: boolean } | undefined { + const overload = isOverloaded(type) ? OverloadedType.getOverloads(type)[0] : type; + if (!overload) { + return undefined; + } + + const name = overload.shared?.name; + if (name !== 'get' && name !== 'pop' && name !== 'setdefault') { + return undefined; + } + + let isBound = false; + let boundType = overload.priv.strippedFirstParamType; + if (boundType) { + isBound = true; + } else if (overload.shared?.parameters && overload.shared.parameters.length > 0) { + boundType = FunctionType.getParamType(overload, 0); + } + + if (boundType && isClassInstance(boundType) && ClassType.isTypedDictClass(boundType)) { + return { classType: boundType, methodName: name, isBound }; + } + + return undefined; +} + +export function applyTypedDictMethodTransform( + evaluator: TypeEvaluator, + errorNode: ExpressionNode, + argList: Arg[], + typedDictClass: ClassType, + methodName: string, + isBound: boolean +): CallResult | undefined { + const keyIndex = isBound ? 0 : 1; + const defaultIndex = isBound ? 1 : 2; + + if (argList.length <= keyIndex) { + return undefined; + } + + const keyArg = argList[keyIndex]; + const keyNode = keyArg.valueExpression ?? errorNode; + const keyTypeResult = keyArg.typeResult ?? evaluator.getTypeOfExpression(keyNode); + const keyType = keyTypeResult.type; + + if (!isUnion(keyType)) { + return undefined; + } + + const defaultArg = argList.length > defaultIndex ? argList[defaultIndex] : undefined; + const defaultNode = defaultArg?.valueExpression ?? errorNode; + const defaultType = defaultArg + ? (defaultArg.typeResult ?? evaluator.getTypeOfExpression(defaultNode)).type + : undefined; + + const entries = getTypedDictMembersForClass(evaluator, typedDictClass, /* allowNarrowed */ methodName === 'get'); + let argumentErrors = false; + + const returnType = mapSubtypes(keyType, (keySubtype) => { + if (isAnyOrUnknown(keySubtype)) { + return keySubtype; + } + + if (isClassInstance(keySubtype) && ClassType.isBuiltIn(keySubtype, 'str')) { + if (keySubtype.priv.literalValue === undefined) { + if (methodName === 'get') { + if (ClassType.isTypedDictEffectivelyClosed(typedDictClass)) { + const extraType = entries.extraItems?.valueType ?? NeverType.createNever(); + return defaultType + ? combineTypes([extraType, defaultType]) + : combineTypes([extraType, evaluator.getNoneType()]); + } + return defaultType + ? combineTypes([AnyType.create(), defaultType]) + : combineTypes([AnyType.create(), evaluator.getNoneType()]); + } else if (methodName === 'pop') { + return defaultType ? combineTypes([UnknownType.create(), defaultType]) : UnknownType.create(); + } else { + return UnknownType.create(); + } + } + + const entryName = keySubtype.priv.literalValue as string; + const entry = entries.knownItems.get(entryName) ?? entries.extraItems; + + if (methodName === 'get') { + if (entry && !isNever(entry.valueType)) { + if (entry.isRequired || entry.isProvided) { + return entry.valueType; + } + return combineTypes([entry.valueType, defaultType ?? evaluator.getNoneType()]); + } + if (ClassType.isTypedDictEffectivelyClosed(typedDictClass)) { + const extraType = entries.extraItems?.valueType; + if (extraType) { + return combineTypes([extraType, defaultType ?? evaluator.getNoneType()]); + } + return defaultType ?? evaluator.getNoneType(); + } + return combineTypes([AnyType.create(), defaultType ?? evaluator.getNoneType()]); + } else if (methodName === 'pop') { + if (entry && !isNever(entry.valueType)) { + if (entry.isReadOnly) { + evaluator.addDiagnostic( + DiagnosticRule.reportTypedDictNotRequiredAccess, + LocAddendum.keyReadOnly().format({ + name: entryName, + type: evaluator.printType(typedDictClass), + }), + keyNode + ); + argumentErrors = true; + return UnknownType.create(); + } + if (entry.isRequired) { + return entry.valueType; + } + return defaultType ? combineTypes([entry.valueType, defaultType]) : entry.valueType; + } + if (defaultType) { + return defaultType; + } + evaluator.addDiagnostic( + DiagnosticRule.reportGeneralTypeIssues, + LocAddendum.keyUndefined().format({ + name: entryName, + type: evaluator.printType(typedDictClass), + }), + keyNode + ); + argumentErrors = true; + return UnknownType.create(); + } else if (methodName === 'setdefault') { + if (entry && !isNever(entry.valueType)) { + if (entry.isReadOnly) { + evaluator.addDiagnostic( + DiagnosticRule.reportGeneralTypeIssues, + LocAddendum.keyReadOnly().format({ + name: entryName, + type: evaluator.printType(typedDictClass), + }), + keyNode + ); + argumentErrors = true; + return UnknownType.create(); + } + if (defaultType && defaultArg) { + const diag = new DiagnosticAddendum(); + if (!evaluator.assignType(entry.valueType, defaultType, diag)) { + evaluator.addDiagnostic( + DiagnosticRule.reportArgumentType, + LocMessage.argAssignmentParam().format({ + paramName: 'default', + paramType: evaluator.printType(entry.valueType), + argType: evaluator.printType(defaultType), + }) + diag.getString(), + defaultNode + ); + argumentErrors = true; + } + } + return entry.valueType; + } + evaluator.addDiagnostic( + DiagnosticRule.reportGeneralTypeIssues, + LocAddendum.keyUndefined().format({ + name: entryName, + type: evaluator.printType(typedDictClass), + }), + keyNode + ); + argumentErrors = true; + return UnknownType.create(); + } + } + + return UnknownType.create(); + }); + + return { + returnType, + argumentErrors, + }; +} + // If the specified type has a non-required key, this method marks the // key as present. export function narrowForKeyAssignment(classType: ClassType, key: string) { diff --git a/packages/pyright-internal/src/tests/samples/typedDict28.py b/packages/pyright-internal/src/tests/samples/typedDict28.py new file mode 100644 index 000000000000..32e3acefe92d --- /dev/null +++ b/packages/pyright-internal/src/tests/samples/typedDict28.py @@ -0,0 +1,45 @@ +# This sample tests type inference and diagnostic behavior for TypedDict +# methods (get, pop, setdefault) when called with union key types. + +from typing import Any, Literal, NotRequired, ReadOnly, TypedDict, assert_type + +class Person(TypedDict): + name: str + age: int + nickname: NotRequired[str] + +class Config(TypedDict): + host: ReadOnly[str] + port: int + +def test_get_union_literal_keys(p: Person, k: Literal["name", "age"]): + v = p.get(k) + assert_type(v, str | int) + +def test_get_union_with_not_required(p: Person, k: Literal["name", "nickname"]): + v = p.get(k) + assert_type(v, str | None) + +def test_get_union_with_unknown_key(p: Person, k: Literal["name", "missing"], default_val: int): + v1 = p.get(k) + assert_type(v1, str | Any | None) + + v2 = p.get(k, default_val) + assert_type(v2, str | Any | int) + +def test_pop_union_literal_keys(p: Person, k: Literal["name", "age"]): + v = p.pop(k) + assert_type(v, str | int) + +def test_pop_readonly_diagnostic(c: Config, k: Literal["host", "port"]): + # This should report an error because "host" is ReadOnly. + c.pop(k) + +def test_setdefault_union_literal_keys(p: Person, k: Literal["name", "age"]): + # This should report an error for default "val" not matching int ("age"). + v = p.setdefault(k, "val") + assert_type(v, str | int) + +def test_unbound_method_union_keys(p: Person, k: Literal["name", "age"]): + v = Person.get(p, k) + assert_type(v, str | int) diff --git a/packages/pyright-internal/src/tests/typeEvaluator7.test.ts b/packages/pyright-internal/src/tests/typeEvaluator7.test.ts index 0776c53fb944..5bd123e3e262 100644 --- a/packages/pyright-internal/src/tests/typeEvaluator7.test.ts +++ b/packages/pyright-internal/src/tests/typeEvaluator7.test.ts @@ -808,6 +808,12 @@ test('TypedDict27', () => { TestUtils.validateResults(analysisResults, 7); }); +test('TypedDict28', () => { + const analysisResults = TestUtils.typeAnalyzeSampleFiles(['typedDict28.py']); + + TestUtils.validateResults(analysisResults, 2); +}); + test('TypedDictInline1', () => { const configOptions = new ConfigOptions(Uri.empty()); configOptions.diagnosticRuleSet.enableExperimentalFeatures = true; From 2b221a1ff8345be0902027f2f4c42abe50167e49 Mon Sep 17 00:00:00 2001 From: Henry Su Date: Mon, 10 Aug 2026 00:20:34 -0500 Subject: [PATCH 2/3] Address review feedback: gate on isSynthesizedMethod, validate receiver and call shape --- .../src/analyzer/typedDicts.ts | 26 ++++++++++++++++--- .../src/tests/samples/typedDict28.py | 24 ++++++++++++++++- .../src/tests/typeEvaluator7.test.ts | 2 +- 3 files changed, 47 insertions(+), 5 deletions(-) diff --git a/packages/pyright-internal/src/analyzer/typedDicts.ts b/packages/pyright-internal/src/analyzer/typedDicts.ts index ce625f95c112..e5296abc47db 100644 --- a/packages/pyright-internal/src/analyzer/typedDicts.ts +++ b/packages/pyright-internal/src/analyzer/typedDicts.ts @@ -1618,7 +1618,7 @@ export function getTypedDictClassFromMethod( type: FunctionType | OverloadedType ): { classType: ClassType; methodName: string; isBound: boolean } | undefined { const overload = isOverloaded(type) ? OverloadedType.getOverloads(type)[0] : type; - if (!overload) { + if (!overload || !FunctionType.isSynthesizedMethod(overload)) { return undefined; } @@ -1650,11 +1650,31 @@ export function applyTypedDictMethodTransform( methodName: string, isBound: boolean ): CallResult | undefined { + // Validate argument counts and shape: + // Bound call: 1 or 2 positional args. Unbound call: 2 or 3 positional args. + const minArgs = isBound ? 1 : 2; + const maxArgs = isBound ? 2 : 3; + if (argList.length < minArgs || argList.length > maxArgs) { + return undefined; + } + + // Require all arguments to be simple positional (no keyword names, no *args/**kwargs) + if (!argList.every((arg) => !arg.name && arg.argCategory === ArgCategory.Simple)) { + return undefined; + } + const keyIndex = isBound ? 0 : 1; const defaultIndex = isBound ? 1 : 2; - if (argList.length <= keyIndex) { - return undefined; + // Validate receiver type for unbound method calls + if (!isBound) { + const selfArg = argList[0]; + const selfNode = selfArg.valueExpression ?? errorNode; + const selfType = (selfArg.typeResult ?? evaluator.getTypeOfExpression(selfNode)).type; + const expectedSelfType = ClassType.cloneAsInstance(typedDictClass); + if (!evaluator.assignType(expectedSelfType, selfType)) { + return undefined; + } } const keyArg = argList[keyIndex]; diff --git a/packages/pyright-internal/src/tests/samples/typedDict28.py b/packages/pyright-internal/src/tests/samples/typedDict28.py index 32e3acefe92d..8bf106f15b02 100644 --- a/packages/pyright-internal/src/tests/samples/typedDict28.py +++ b/packages/pyright-internal/src/tests/samples/typedDict28.py @@ -1,7 +1,7 @@ # This sample tests type inference and diagnostic behavior for TypedDict # methods (get, pop, setdefault) when called with union key types. -from typing import Any, Literal, NotRequired, ReadOnly, TypedDict, assert_type +from typing import Any, Literal, NotRequired, ReadOnly, TypedDict, assert_type, overload class Person(TypedDict): name: str @@ -43,3 +43,25 @@ def test_setdefault_union_literal_keys(p: Person, k: Literal["name", "age"]): def test_unbound_method_union_keys(p: Person, k: Literal["name", "age"]): v = Person.get(p, k) assert_type(v, str | int) + +class CustomContainer: + @overload + def get(self, key: Literal["a"]) -> int: ... + @overload + def get(self, key: Literal["b"]) -> str: ... + def get(self, key: str) -> Any: + pass + +def test_custom_overloaded_get_not_intercepted(c: CustomContainer, k: Literal["a", "b"]): + # User-defined overloaded function should NOT be intercepted by TypedDict transform. + # Standard overload resolution should apply. + v = c.get(k) + +def test_unbound_invalid_receiver(k: Literal["name", "age"]): + # Unbound call with invalid receiver should fall back to normal overload validation. + Person.get(123, k) + +def test_keyword_and_extra_args(p: Person, k: Literal["name", "age"]): + # Keyword arguments or extra arguments fall back to standard overload validation. + p.get(k, default=0) + p.get(k, 0, 1) diff --git a/packages/pyright-internal/src/tests/typeEvaluator7.test.ts b/packages/pyright-internal/src/tests/typeEvaluator7.test.ts index 5bd123e3e262..ff5ba8e6dde6 100644 --- a/packages/pyright-internal/src/tests/typeEvaluator7.test.ts +++ b/packages/pyright-internal/src/tests/typeEvaluator7.test.ts @@ -811,7 +811,7 @@ test('TypedDict27', () => { test('TypedDict28', () => { const analysisResults = TestUtils.typeAnalyzeSampleFiles(['typedDict28.py']); - TestUtils.validateResults(analysisResults, 2); + TestUtils.validateResults(analysisResults, 5); }); test('TypedDictInline1', () => { From b4b2d232be16a0fd9e95dd98924db770e77e7e0d Mon Sep 17 00:00:00 2001 From: Henry Su Date: Mon, 10 Aug 2026 11:42:28 -0500 Subject: [PATCH 3/3] Address review feedback: use isMethodType, enforce default arg for setdefault, validate non-string key subtypes, propagate isTypeIncomplete --- .../src/analyzer/typedDicts.ts | 41 ++++++++++++++++--- .../src/tests/samples/typedDict28.py | 8 ++++ .../src/tests/typeEvaluator7.test.ts | 2 +- 3 files changed, 44 insertions(+), 7 deletions(-) diff --git a/packages/pyright-internal/src/analyzer/typedDicts.ts b/packages/pyright-internal/src/analyzer/typedDicts.ts index e5296abc47db..09632ec192f4 100644 --- a/packages/pyright-internal/src/analyzer/typedDicts.ts +++ b/packages/pyright-internal/src/analyzer/typedDicts.ts @@ -53,6 +53,7 @@ import { isClass, isClassInstance, isInstantiableClass, + isMethodType, isNever, isOverloaded, isUnion, @@ -1627,10 +1628,10 @@ export function getTypedDictClassFromMethod( return undefined; } - let isBound = false; - let boundType = overload.priv.strippedFirstParamType; - if (boundType) { - isBound = true; + const isBound = isMethodType(overload); + let boundType: Type | undefined; + if (isBound) { + boundType = overload.priv.strippedFirstParamType; } else if (overload.shared?.parameters && overload.shared.parameters.length > 0) { boundType = FunctionType.getParamType(overload, 0); } @@ -1651,8 +1652,9 @@ export function applyTypedDictMethodTransform( isBound: boolean ): CallResult | undefined { // Validate argument counts and shape: - // Bound call: 1 or 2 positional args. Unbound call: 2 or 3 positional args. - const minArgs = isBound ? 1 : 2; + // Bound call: 1 or 2 positional args for get/pop, 2 for setdefault. + // Unbound call: 2 or 3 positional args for get/pop, 3 for setdefault. + const minArgs = isBound ? (methodName === 'setdefault' ? 2 : 1) : methodName === 'setdefault' ? 3 : 2; const maxArgs = isBound ? 2 : 3; if (argList.length < minArgs || argList.length > maxArgs) { return undefined; @@ -1686,7 +1688,33 @@ export function applyTypedDictMethodTransform( return undefined; } + // Verify all key subtypes are valid string types. If any subtype is not assignable to str, + // fall back to standard overload validation so argument type errors are reported. + const strType = evaluator.getBuiltInObject(errorNode, 'str'); + let hasInvalidKeySubtype = false; + mapSubtypes(keyType, (keySubtype) => { + if (isAnyOrUnknown(keySubtype)) { + return keySubtype; + } + if (!evaluator.assignType(strType, keySubtype)) { + hasInvalidKeySubtype = true; + } + return keySubtype; + }); + + if (hasInvalidKeySubtype) { + return undefined; + } + + let isTypeIncomplete = !!keyTypeResult.isIncomplete; const defaultArg = argList.length > defaultIndex ? argList[defaultIndex] : undefined; + if (defaultArg?.typeResult?.isIncomplete) { + isTypeIncomplete = true; + } + if (!isBound && argList[0].typeResult?.isIncomplete) { + isTypeIncomplete = true; + } + const defaultNode = defaultArg?.valueExpression ?? errorNode; const defaultType = defaultArg ? (defaultArg.typeResult ?? evaluator.getTypeOfExpression(defaultNode)).type @@ -1819,6 +1847,7 @@ export function applyTypedDictMethodTransform( return { returnType, argumentErrors, + isTypeIncomplete, }; } diff --git a/packages/pyright-internal/src/tests/samples/typedDict28.py b/packages/pyright-internal/src/tests/samples/typedDict28.py index 8bf106f15b02..baf03ec14170 100644 --- a/packages/pyright-internal/src/tests/samples/typedDict28.py +++ b/packages/pyright-internal/src/tests/samples/typedDict28.py @@ -65,3 +65,11 @@ def test_keyword_and_extra_args(p: Person, k: Literal["name", "age"]): # Keyword arguments or extra arguments fall back to standard overload validation. p.get(k, default=0) p.get(k, 0, 1) + +def test_setdefault_missing_default(p: Person, k: Literal["name", "age"]): + # setdefault requires default argument; missing default falls back to normal validation and errors. + p.setdefault(k) + +def test_union_with_non_string_subtype(p: Person, k: Literal["name"] | int): + # Union containing non-string subtype falls back to normal validation and errors on int. + p.get(k) diff --git a/packages/pyright-internal/src/tests/typeEvaluator7.test.ts b/packages/pyright-internal/src/tests/typeEvaluator7.test.ts index ff5ba8e6dde6..2d6a78c6178c 100644 --- a/packages/pyright-internal/src/tests/typeEvaluator7.test.ts +++ b/packages/pyright-internal/src/tests/typeEvaluator7.test.ts @@ -811,7 +811,7 @@ test('TypedDict27', () => { test('TypedDict28', () => { const analysisResults = TestUtils.typeAnalyzeSampleFiles(['typedDict28.py']); - TestUtils.validateResults(analysisResults, 5); + TestUtils.validateResults(analysisResults, 8); }); test('TypedDictInline1', () => {