diff --git a/mypy_django_plugin/django/context.py b/mypy_django_plugin/django/context.py index 73376c720..461c3a341 100644 --- a/mypy_django_plugin/django/context.py +++ b/mypy_django_plugin/django/context.py @@ -5,7 +5,7 @@ from collections import defaultdict from contextlib import contextmanager from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import TYPE_CHECKING, Any, Literal, NamedTuple from django.core.exceptions import FieldDoesNotExist, FieldError from django.db import models @@ -111,6 +111,14 @@ def _get_field_get_type_from_model_type_info(info: TypeInfo | None, field_name: return None +class ResolvedLookupField(NamedTuple): + """Result of resolving a lookup's field parts to a concrete field.""" + + field: Field[Any, Any] | ForeignObjectRel + model: type[Model] + is_nullable: bool + + class DjangoContext: def __init__(self, django_settings_module: str) -> None: self.django_settings_module = django_settings_module @@ -158,7 +166,13 @@ def get_model_relations(self, model_cls: type[Model]) -> Iterator[ForeignObjectR if isinstance(field, ForeignObjectRel): yield field - def get_field_lookup_exact_type(self, api: TypeChecker, field: Field[Any, Any] | ForeignObjectRel) -> MypyType: + def get_field_lookup_exact_type( + self, + api: TypeChecker, + field: Field[Any, Any] | ForeignObjectRel, + *, + is_nullable: bool = False, + ) -> MypyType: if isinstance(field, RelatedField | ForeignObjectRel): related_model_cls = self.get_field_related_model_cls(field) rel_model_info = helpers.lookup_class_typeinfo(api, related_model_cls) @@ -174,7 +188,9 @@ def get_field_lookup_exact_type(self, api: TypeChecker, field: Field[Any, Any] | field_info = helpers.lookup_class_typeinfo(api, field.__class__) if field_info is None: return AnyType(TypeOfAny.explicit) - return helpers.get_private_descriptor_type(field_info, "_pyi_lookup_exact_type", is_nullable=field.null) + return helpers.get_private_descriptor_type( + field_info, "_pyi_lookup_exact_type", is_nullable=field.null or is_nullable + ) def get_related_target_field( self, related_model_cls: type[Model], field: ForeignKey[Any, Any] @@ -395,11 +411,19 @@ def get_field_related_model_cls(self, field: RelatedField[Any, Any] | ForeignObj return related_model_cls - def _resolve_field_from_parts( - self, field_parts: Iterable[str], model_cls: type[Model] - ) -> tuple[Field[Any, Any] | ForeignObjectRel, type[Model]]: + def _resolve_field_from_parts(self, field_parts: Iterable[str], model_cls: type[Model]) -> ResolvedLookupField: + """ + Resolve ``field_parts`` (as produced by ``Query.solve_lookup_type``) to a + concrete field and the model it belongs to. ``is_nullable`` is ``True`` + when the resolved field should be treated as nullable even if + ``field.null`` is ``False`` — this happens when the lookup goes through + a nullable FK ``_id`` attname (e.g. ``prospect_id`` on + ``prospect = ForeignKey(null=True)``), because the underlying PK field + is always ``null=False`` but the column is nullable when the FK is. + """ currently_observed_model = model_cls field: _AnyField | None = None + is_nullable = False for field_part in field_parts: if field_part == "pk": field = self.get_primary_key_field(currently_observed_model) @@ -408,8 +432,17 @@ def _resolve_field_from_parts( field = currently_observed_model._meta.get_field(field_part) if isinstance(field, RelatedField): currently_observed_model = self.get_field_related_model_cls(field) - model_name = currently_observed_model._meta.model_name - if model_name is not None and field_part == (model_name + "_id"): + if field_part == field.attname and field_part != field.name: + # The lookup uses the FK's ``_id`` suffix (e.g. ``prospect_id``). + # Django's ``get_field`` returns the ``RelatedField`` for this + # attname, but semantically it refers to the underlying column, + # so resolve to the related model's primary key field instead. + # Preserve the original FK's nullability — the PK field itself + # is always ``null=False`` but the column is nullable when the + # FK is. + # ``field_part != field.name`` excludes ManyToManyField, whose + # ``attname`` equals its ``name`` (no ``_id`` column variant). + is_nullable = field.null field = self.get_primary_key_field(currently_observed_model) if isinstance(field, ForeignObjectRel): @@ -417,7 +450,7 @@ def _resolve_field_from_parts( # Guaranteed by `query.solve_lookup_type` before. assert isinstance(field, Field | ForeignObjectRel) - return field, currently_observed_model + return ResolvedLookupField(field, currently_observed_model, is_nullable) def solve_lookup_type( self, model_cls: type[Model], lookup: str @@ -462,7 +495,8 @@ def resolve_lookup_into_field( lookup_parts, field_parts, _ = solved_lookup if lookup_parts: raise LookupsAreUnsupported() - return self._resolve_field_from_parts(field_parts, model_cls) + field, model, _ = self._resolve_field_from_parts(field_parts, model_cls) + return field, model def _resolve_lookup_type_from_lookup_class( self, ctx: MethodContext, lookup_cls: type, field: Field[Any, Any] | ForeignObjectRel | None = None @@ -557,7 +591,9 @@ def resolve_lookup_expected_type( if is_expression: return AnyType(TypeOfAny.explicit) - field, _ = self._resolve_field_from_parts(field_parts, model_cls) + resolved = self._resolve_field_from_parts(field_parts, model_cls) + field = resolved.field + is_nullable = resolved.is_nullable lookup_cls = None if lookup_parts: @@ -568,10 +604,12 @@ def resolve_lookup_expected_type( return AnyType(TypeOfAny.explicit) if lookup_cls is None or issubclass(lookup_cls, Exact): - return self.get_field_lookup_exact_type(helpers.get_typechecker_api(ctx), field) + return self.get_field_lookup_exact_type(helpers.get_typechecker_api(ctx), field, is_nullable=is_nullable) if issubclass(lookup_cls, In): - exact_type = self.get_field_lookup_exact_type(helpers.get_typechecker_api(ctx), field) + exact_type = self.get_field_lookup_exact_type( + helpers.get_typechecker_api(ctx), field, is_nullable=is_nullable + ) return ctx.api.named_generic_type("typing.Iterable", [exact_type]) resolved_type = self._resolve_lookup_type_from_lookup_class(ctx, lookup_cls, field) diff --git a/tests/typecheck/managers/querysets/test_filter.yml b/tests/typecheck/managers/querysets/test_filter.yml index 6ee720991..86cd5c03b 100644 --- a/tests/typecheck/managers/querysets/test_filter.yml +++ b/tests/typecheck/managers/querysets/test_filter.yml @@ -177,6 +177,57 @@ publisher = models.ForeignKey(Publisher, on_delete=models.CASCADE, related_name='blogs') +- case: nullable_foreign_key_id_lookup + main: | + from myapp.models import Article, Author + author = Author() + + # Nullable FK: filtering on the _id column with None is valid + Article.objects.filter(author_id=None) + Article.objects.filter(author_id=1) + Article.objects.filter(author_id=author) # E: Incompatible type for lookup 'author_id': (got "Author", expected "str | int | None") [misc] + + # Non-nullable FK: None is not valid for the _id column + Article.objects.filter(blog_id=None) # E: Incompatible type for lookup 'blog_id': (got "None", expected "str | int") [misc] + Article.objects.filter(blog_id=1) + installed_apps: + - myapp + files: + - path: myapp/__init__.py + - path: myapp/models.py + content: | + from django.db import models + class Author(models.Model): + pass + class Blog(models.Model): + pass + class Article(models.Model): + author = models.ForeignKey(Author, null=True, on_delete=models.SET_NULL, related_name='articles') + blog = models.ForeignKey(Blog, on_delete=models.CASCADE, related_name='articles') + + +- case: nullable_one_to_one_id_lookup + main: | + from myapp.models import User, Profile + user = User() + + # Nullable O2O: filtering on the _id column with None is valid + Profile.objects.filter(user_id=None) + Profile.objects.filter(user_id=1) + Profile.objects.filter(user_id=user) # E: Incompatible type for lookup 'user_id': (got "User", expected "str | int | None") [misc] + installed_apps: + - myapp + files: + - path: myapp/__init__.py + - path: myapp/models.py + content: | + from django.db import models + class User(models.Model): + pass + class Profile(models.Model): + user = models.OneToOneField(User, null=True, on_delete=models.SET_NULL) + + - case: related_model_reverse_foreign_key_lookup main: | from myapp.models import Blog, Publisher, Category