diff --git a/packages/pyright-internal/src/analyzer/patternMatching.ts b/packages/pyright-internal/src/analyzer/patternMatching.ts index 62847394d58a..928e3d5ce910 100644 --- a/packages/pyright-internal/src/analyzer/patternMatching.ts +++ b/packages/pyright-internal/src/analyzer/patternMatching.ts @@ -1334,6 +1334,31 @@ function narrowTypeBasedOnValuePattern( (subjectSubtypeExpanded) => { // If this is a negative test, see if it's an enum value. if (!isPositiveTest) { + if ( + isInstantiableClass(subjectSubtypeExpanded) && + isInstantiableClass(valueSubtypeExpanded) && + isSameWithoutLiteralValue(subjectSubtypeExpanded, valueSubtypeExpanded) + ) { + const metaclass = subjectSubtypeExpanded.shared.effectiveMetaclass; + let eqClass: ClassType | undefined; + if (metaclass && isInstantiableClass(metaclass)) { + const eqMember = lookUpClassMember(metaclass, '__eq__'); + if (eqMember && isClass(eqMember.classType)) { + eqClass = eqMember.classType; + } + } + + const isStandardEquality = + !eqClass || ClassType.isBuiltIn(eqClass, ['type', 'object', 'ABCMeta', 'EnumMeta']); + + if ( + isStandardEquality && + (ClassType.isFinal(subjectSubtypeExpanded) || + !subjectSubtypeExpanded.priv.includeSubclasses) + ) { + return undefined; + } + } if ( isClassInstance(subjectSubtypeExpanded) && isClassInstance(valueSubtypeExpanded) && diff --git a/packages/pyright-internal/src/tests/samples/matchClassFinal.py b/packages/pyright-internal/src/tests/samples/matchClassFinal.py new file mode 100644 index 000000000000..786aca8c5be4 --- /dev/null +++ b/packages/pyright-internal/src/tests/samples/matchClassFinal.py @@ -0,0 +1,87 @@ +# pyright: reportMatchNotExhaustive=true + +from typing import final + +class Base: + pass + +@final +class A1(Base): + pass + +@final +class B1(Base): + pass + +class NS1: + A1 = A1 + B1 = B1 + +def exhaustive_final(inst: A1 | B1): + match type(inst): + case NS1.A1: + pass + case NS1.B1: + pass + +class A2(Base): + pass + +class B2(Base): + pass + +class NS2: + A2 = A2 + B2 = B2 + +def non_exhaustive_non_final(inst: A2 | B2): + match type(inst): + case NS2.A2: + pass + case NS2.B2: + pass + +class Meta(type): + def __eq__(self, other): + return False + +@final +class C1(metaclass=Meta): + pass + +@final +class D1(metaclass=Meta): + pass + +class NS3: + C1 = C1 + D1 = D1 + +def non_exhaustive_custom_meta(inst: C1 | D1): + match type(inst): + case NS3.C1: + pass + case NS3.D1: + pass + +class MetaNoOverride(type): + pass + +@final +class E1(metaclass=MetaNoOverride): + pass + +@final +class F1(metaclass=MetaNoOverride): + pass + +class NS4: + E1 = E1 + F1 = F1 + +def exhaustive_custom_meta_no_override(inst: E1 | F1): + match type(inst): + case NS4.E1: + pass + case NS4.F1: + pass diff --git a/packages/pyright-internal/src/tests/typeEvaluator6.test.ts b/packages/pyright-internal/src/tests/typeEvaluator6.test.ts index 652e0c1479f7..51e5a6b51709 100644 --- a/packages/pyright-internal/src/tests/typeEvaluator6.test.ts +++ b/packages/pyright-internal/src/tests/typeEvaluator6.test.ts @@ -623,6 +623,14 @@ test('MatchMapping1', () => { TestUtils.validateResults(analysisResults, 2); }); +test('MatchClassFinal', () => { + const configOptions = new ConfigOptions(Uri.empty()); + + configOptions.defaultPythonVersion = pythonVersion3_12; + const analysisResults = TestUtils.typeAnalyzeSampleFiles(['matchClassFinal.py'], configOptions); + TestUtils.validateResults(analysisResults, 2); // 1 error for non-final, 1 error for custom metaclass +}); + test('MatchLiteral1', () => { const configOptions = new ConfigOptions(Uri.empty());