mirror of
https://github.com/davidhalter/django-stubs.git
synced 2025-12-06 20:24:31 +08:00
225 lines
9.3 KiB
Python
225 lines
9.3 KiB
Python
import os
|
|
from collections import defaultdict
|
|
from contextlib import contextmanager
|
|
from typing import Any, Dict, Iterator, List, Optional, TYPE_CHECKING, Tuple, Type
|
|
|
|
from django.core.exceptions import FieldError
|
|
from django.db.models.base import Model
|
|
from django.db.models.fields.related import ForeignKey, RelatedField
|
|
from django.db.models.fields.reverse_related import ForeignObjectRel
|
|
from django.db.models.sql.query import Query
|
|
from django.utils.functional import cached_property
|
|
from mypy.checker import TypeChecker
|
|
from mypy.types import Instance, Type as MypyType
|
|
|
|
from django.contrib.postgres.fields import ArrayField
|
|
from django.db.models.fields import CharField, Field
|
|
from mypy_django_plugin.lib import helpers
|
|
|
|
if TYPE_CHECKING:
|
|
from django.apps.registry import Apps
|
|
from django.conf import LazySettings
|
|
|
|
|
|
@contextmanager
|
|
def temp_environ():
|
|
"""Allow the ability to set os.environ temporarily"""
|
|
environ = dict(os.environ)
|
|
try:
|
|
yield
|
|
finally:
|
|
os.environ.clear()
|
|
os.environ.update(environ)
|
|
|
|
|
|
def initialize_django(settings_module: str) -> Tuple['Apps', 'LazySettings']:
|
|
with temp_environ():
|
|
os.environ['DJANGO_SETTINGS_MODULE'] = settings_module
|
|
|
|
def noop_class_getitem(cls, key):
|
|
return cls
|
|
|
|
from django.db import models
|
|
|
|
models.QuerySet.__class_getitem__ = classmethod(noop_class_getitem)
|
|
models.Manager.__class_getitem__ = classmethod(noop_class_getitem)
|
|
|
|
from django.conf import settings
|
|
from django.apps import apps
|
|
|
|
apps.get_models.cache_clear()
|
|
apps.get_swappable_settings_name.cache_clear()
|
|
|
|
if not settings.configured:
|
|
settings._setup()
|
|
|
|
apps.populate(settings.INSTALLED_APPS)
|
|
|
|
assert apps.apps_ready
|
|
assert settings.configured
|
|
|
|
return apps, settings
|
|
|
|
|
|
class DjangoFieldsContext:
|
|
def __init__(self, django_context: 'DjangoContext') -> None:
|
|
self.django_context = django_context
|
|
|
|
def get_attname(self, field: Field) -> str:
|
|
attname = field.attname
|
|
return attname
|
|
|
|
def get_field_nullability(self, field: Field, method: Optional[str]) -> bool:
|
|
nullable = field.null
|
|
if not nullable and isinstance(field, CharField) and field.blank:
|
|
return True
|
|
if method == '__init__':
|
|
if field.primary_key or isinstance(field, ForeignKey):
|
|
return True
|
|
if field.has_default():
|
|
return True
|
|
return nullable
|
|
|
|
def get_field_set_type(self, api: TypeChecker, field: Field, method: str) -> MypyType:
|
|
target_field = field
|
|
if isinstance(field, ForeignKey):
|
|
target_field = field.target_field
|
|
|
|
field_info = helpers.lookup_class_typeinfo(api, target_field.__class__)
|
|
field_set_type = helpers.get_private_descriptor_type(field_info, '_pyi_private_set_type',
|
|
is_nullable=self.get_field_nullability(field, method))
|
|
if isinstance(target_field, ArrayField):
|
|
argument_field_type = self.get_field_set_type(api, target_field.base_field, method)
|
|
field_set_type = helpers.convert_any_to_type(field_set_type, argument_field_type)
|
|
return field_set_type
|
|
|
|
def get_field_get_type(self, api: TypeChecker, field: Field, method: str) -> MypyType:
|
|
field_info = helpers.lookup_class_typeinfo(api, field.__class__)
|
|
is_nullable = self.get_field_nullability(field, method)
|
|
if isinstance(field, RelatedField):
|
|
if method == 'values':
|
|
primary_key_field = self.django_context.get_primary_key_field(field.related_model)
|
|
return self.get_field_get_type(api, primary_key_field, method)
|
|
|
|
model_info = helpers.lookup_class_typeinfo(api, field.related_model)
|
|
return Instance(model_info, [])
|
|
else:
|
|
return helpers.get_private_descriptor_type(field_info, '_pyi_private_get_type',
|
|
is_nullable=is_nullable)
|
|
|
|
|
|
class DjangoLookupsContext:
|
|
def __init__(self, django_context: 'DjangoContext'):
|
|
self.django_context = django_context
|
|
|
|
def resolve_lookup(self, model_cls: Type[Model], lookup: str) -> Field:
|
|
query = Query(model_cls)
|
|
lookup_parts, field_parts, is_expression = query.solve_lookup_type(lookup)
|
|
if lookup_parts:
|
|
raise FieldError('Lookups not supported yet')
|
|
|
|
currently_observed_model = model_cls
|
|
current_field = None
|
|
for field_part in field_parts:
|
|
if field_part == 'pk':
|
|
return self.django_context.get_primary_key_field(currently_observed_model)
|
|
|
|
current_field = currently_observed_model._meta.get_field(field_part)
|
|
if isinstance(current_field, ForeignObjectRel):
|
|
currently_observed_model = current_field.related_model
|
|
current_field = self.django_context.get_primary_key_field(currently_observed_model)
|
|
else:
|
|
if isinstance(current_field, RelatedField):
|
|
currently_observed_model = current_field.related_model
|
|
|
|
return current_field
|
|
|
|
|
|
class DjangoContext:
|
|
def __init__(self, plugin_toml_config: Optional[Dict[str, Any]]) -> None:
|
|
self.fields_context = DjangoFieldsContext(self)
|
|
self.lookups_context = DjangoLookupsContext(self)
|
|
|
|
self.django_settings_module = None
|
|
if plugin_toml_config:
|
|
self.django_settings_module = plugin_toml_config.get('django_settings_module', None)
|
|
|
|
self.apps_registry: Optional[Dict[str, str]] = None
|
|
self.settings: LazySettings = None
|
|
if self.django_settings_module:
|
|
apps, settings = initialize_django(self.django_settings_module)
|
|
self.apps_registry = apps
|
|
self.settings = settings
|
|
|
|
@cached_property
|
|
def model_modules(self) -> Dict[str, List[Type[Model]]]:
|
|
""" All modules that contain Django models. """
|
|
if self.apps_registry is None:
|
|
return {}
|
|
|
|
modules: Dict[str, List[Type[Model]]] = defaultdict(list)
|
|
for model_cls in self.apps_registry.get_models():
|
|
modules[model_cls.__module__].append(model_cls)
|
|
return modules
|
|
|
|
def get_model_class_by_fullname(self, fullname: str) -> Optional[Type[Model]]:
|
|
# Returns None if Model is abstract
|
|
module, _, model_cls_name = fullname.rpartition('.')
|
|
for model_cls in self.model_modules.get(module, []):
|
|
if model_cls.__name__ == model_cls_name:
|
|
return model_cls
|
|
|
|
def get_model_fields(self, model_cls: Type[Model]) -> Iterator[Field]:
|
|
for field in model_cls._meta.get_fields():
|
|
if isinstance(field, Field):
|
|
yield field
|
|
|
|
def get_model_relations(self, model_cls: Type[Model]) -> Iterator[ForeignObjectRel]:
|
|
for field in model_cls._meta.get_fields():
|
|
if isinstance(field, ForeignObjectRel):
|
|
yield field
|
|
|
|
def get_primary_key_field(self, model_cls: Type[Model]) -> Field:
|
|
for field in model_cls._meta.get_fields():
|
|
if isinstance(field, Field):
|
|
if field.primary_key:
|
|
return field
|
|
raise ValueError('No primary key defined')
|
|
|
|
def get_expected_types(self, api: TypeChecker, model_cls: Type[Model], method: str) -> Dict[str, MypyType]:
|
|
from django.contrib.contenttypes.fields import GenericForeignKey
|
|
|
|
expected_types = {}
|
|
# add pk
|
|
primary_key_field = self.get_primary_key_field(model_cls)
|
|
field_set_type = self.fields_context.get_field_set_type(api, primary_key_field, method)
|
|
expected_types['pk'] = field_set_type
|
|
|
|
for field in model_cls._meta.get_fields():
|
|
if isinstance(field, Field):
|
|
field_name = field.attname
|
|
field_set_type = self.fields_context.get_field_set_type(api, field, method)
|
|
expected_types[field_name] = field_set_type
|
|
|
|
if isinstance(field, ForeignKey):
|
|
field_name = field.name
|
|
foreign_key_info = helpers.lookup_class_typeinfo(api, field.__class__)
|
|
related_model_info = helpers.lookup_class_typeinfo(api, field.related_model)
|
|
is_nullable = self.fields_context.get_field_nullability(field, method)
|
|
foreign_key_set_type = helpers.get_private_descriptor_type(foreign_key_info,
|
|
'_pyi_private_set_type',
|
|
is_nullable=is_nullable)
|
|
model_set_type = helpers.convert_any_to_type(foreign_key_set_type,
|
|
Instance(related_model_info, []))
|
|
expected_types[field_name] = model_set_type
|
|
|
|
elif isinstance(field, GenericForeignKey):
|
|
# it's generic, so cannot set specific model
|
|
field_name = field.name
|
|
gfk_info = helpers.lookup_class_typeinfo(api, field.__class__)
|
|
gfk_set_type = helpers.get_private_descriptor_type(gfk_info, '_pyi_private_set_type',
|
|
is_nullable=True)
|
|
expected_types[field_name] = gfk_set_type
|
|
|
|
return expected_types
|