修复了 QModel 链式表存在的一些问题

This commit is contained in:
2026-06-16 23:33:23 +08:00
parent 134f7e2ce0
commit 54a64dc220
2 changed files with 269 additions and 83 deletions

View File

@@ -1,13 +1,22 @@
from __future__ import annotations
import threading
import functools
from typing import Any, Dict, Iterator, List, Optional, Tuple, Type, Union, overload
from django.db import connection as default_db_connection, models, transaction
from django.db.models import Q, F, Count, Sum, Avg, Min, Max, Subquery, OuterRef, Case, When, Value
from gvsdsdk.model_base import FieldExpression
_local = threading.local()
def _get_field_expression_cls():
"""延迟导入,避免 fluent ↔ model_base 循环依赖"""
from gvsdsdk.model_base import FieldExpression
return FieldExpression
class _AggregateExpr:
def __init__(self, model, aggregate):
self.model = model
@@ -140,7 +149,7 @@ class RelationAccessor:
try:
self._model._meta.get_field(name)
full_name = f'{self._prefix}__{name}' if self._prefix else name
return FieldExpression(self._model, full_name)
return _get_field_expression_cls()(self._model, full_name)
except Exception:
for rel in self._model._meta.get_fields(include_hidden=True):
if hasattr(rel, 'related_model') and rel.related_model:
@@ -152,7 +161,7 @@ class RelationAccessor:
if rel_name and rel_name == name:
full_prefix = f'{self._prefix}__{rel_name}' if self._prefix else rel_name
return RelationAccessor(rel.related_model, full_prefix)
return FieldExpression(self._model, name)
return _get_field_expression_cls()(self._model, name)
class _ValuesResult:
@@ -162,38 +171,43 @@ class _ValuesResult:
额外提供 .all() 返回字典列表,.first() 返回字典或 None
.scalar() 返回单字段场景的标量值。
"""
def __init__(self, qs, fields, original_fields):
self._qs = qs # Django ValuesQuerySet返回字典
self._fields = fields # 转换后的字段名列表(传给 Django .values() 的)
self._original_fields = original_fields # 原始传入的字段(用于判断单字段场景)
def __init__(self, qs: models.QuerySet, fields: List[str], original_fields: Tuple) -> None:
self._qs: models.QuerySet = qs # Django ValuesQuerySet返回字典
self._fields: List[str] = fields # 转换后的字段名列表(传给 Django .values() 的)
self._original_fields: Tuple = original_fields # 原始传入的字段(用于判断单字段场景)
def filter(self, *args, **kwargs):
@property
def query(self) -> Any:
"""暴露底层 Django QuerySet 的 query 属性,支持 Subquery() 等场景。"""
return self._qs.query
def filter(self, *args: Any, **kwargs: Any) -> _ValuesResult:
self._qs = self._qs.filter(*args, **kwargs)
return self
def exclude(self, *args, **kwargs):
def exclude(self, *args: Any, **kwargs: Any) -> _ValuesResult:
self._qs = self._qs.exclude(*args, **kwargs)
return self
def order_by(self, *fields):
def order_by(self, *fields: str) -> _ValuesResult:
self._qs = self._qs.order_by(*fields)
return self
def distinct(self):
def distinct(self) -> _ValuesResult:
self._qs = self._qs.distinct()
return self
def annotate(self, **kwargs):
def annotate(self, **kwargs: Any) -> _ValuesResult:
self._qs = self._qs.annotate(**kwargs)
return self
def first(self):
def first(self) -> Optional[Dict[str, Any]]:
return self._qs.first()
def all(self):
def all(self) -> List[Dict[str, Any]]:
return list(self._qs)
def scalar(self):
def scalar(self) -> Any:
row = self._qs.first()
if row is None:
return None
@@ -201,52 +215,80 @@ class _ValuesResult:
return row.get(self._fields[0])
return row
def count(self):
def count(self) -> int:
return self._qs.count()
def exists(self):
def exists(self) -> bool:
return self._qs.exists()
def __getitem__(self, key):
def __getitem__(self, key: Union[int, slice]) -> Any:
if isinstance(key, slice):
self._qs = self._qs[key]
return self
else:
return self._qs[key]
def __iter__(self):
def __iter__(self) -> Iterator[Dict[str, Any]]:
return iter(self._qs)
def __len__(self):
def __len__(self) -> int:
return self.count()
def __bool__(self):
def __bool__(self) -> bool:
return self._qs.exists()
class FluentQuery:
def __init__(self, model_class):
self._model = model_class
self._qs = model_class.objects.all()
self._joined_models = []
self._cache_enabled = True
"""QModel 兼容层查询构建器,提供 SQLAlchemy 风格的链式 OOP 查询接口。
def _clone(self):
用法示例::
# 基础查询
Dingdan.query.filter(zhuangtai=8).all()
# 链式过滤
Dingdan.query.filter(jine__gt=100).order_by('-create_time').limit(10)
# 字段表达式(通过 Model.field 自动获取 FieldExpression
Dingdan.query.filter(Dingdan.jine > 100, Dingdan.zhuangtai.in_([1,2,3]))
# 聚合
Dingdan.query.aggregate(total=Sum('jine'))
# values / annotate
Dingdan.query.values('zhuangtai').annotate(cnt=Count('id'))
# 分页
Dingdan.query.paginate(page=1, per_page=20)
"""
def __init__(self, model_class: Type[models.Model]) -> None:
self._model: Type[models.Model] = model_class
self._qs: models.QuerySet = model_class.objects.all()
self._joined_models: List[Type[models.Model]] = []
self._cache_enabled: bool = True
@property
def query(self) -> Any:
"""暴露底层 Django QuerySet 的 query 属性,支持 Subquery() 等场景。"""
return self._qs.query
def _clone(self) -> FluentQuery:
q = FluentQuery(self._model)
q._qs = self._qs
q._joined_models = list(self._joined_models)
q._cache_enabled = self._cache_enabled
return q
def _cache_key(self, method):
def _cache_key(self, method: str) -> Tuple[Any, str, Any]:
try:
query_str = str(self._qs.query)
except Exception:
query_str = id(self._qs)
return (self._model, method, query_str)
def _remap_q(self, q):
new_children = []
def _remap_q(self, q: Q) -> Q:
new_children: List[Any] = []
for child in q.children:
if isinstance(child, Q):
new_children.append(self._remap_q(child))
@@ -263,7 +305,7 @@ class FluentQuery:
return new_q
return q
def _remap_key(self, key):
def _remap_key(self, key: str) -> str:
parts = key.split('__')
field_name = parts[0]
try:
@@ -288,8 +330,9 @@ class FluentQuery:
pass
return key
def filter(self, *args, **kwargs):
converted_args = []
def filter(self, *args: Any, **kwargs: Any) -> FluentQuery:
"""过滤查询,支持 Q 对象和关键字参数,返回 self 以支持链式调用。"""
converted_args: List[Any] = []
for arg in args:
if isinstance(arg, Q) and self._joined_models:
converted_args.append(self._remap_q(arg))
@@ -298,8 +341,9 @@ class FluentQuery:
self._qs = self._qs.filter(*converted_args, **kwargs)
return self
def exclude(self, *args, **kwargs):
converted_args = []
def exclude(self, *args: Any, **kwargs: Any) -> FluentQuery:
"""排除查询,返回 self 以支持链式调用。"""
converted_args: List[Any] = []
for arg in args:
if isinstance(arg, Q) and self._joined_models:
converted_args.append(self._remap_q(arg))
@@ -308,19 +352,21 @@ class FluentQuery:
self._qs = self._qs.exclude(*converted_args, **kwargs)
return self
def order_by(self, *args):
converted = []
def order_by(self, *args: Any) -> FluentQuery:
"""排序,支持字符串和 FieldExpression返回 self 以支持链式调用。"""
converted: List[str] = []
for a in args:
if isinstance(a, str):
converted.append(a)
elif isinstance(a, FieldExpression):
elif isinstance(a, _get_field_expression_cls()):
converted.append(str(a))
else:
converted.append(str(a))
self._qs = self._qs.order_by(*converted)
return self
def first(self):
def first(self) -> Optional[models.Model]:
"""返回第一条记录或 None。"""
if self._cache_enabled:
cache = _get_cache()
key = self._cache_key('first')
@@ -333,13 +379,13 @@ class FluentQuery:
_get_cache()[self._cache_key('first')] = obj
return obj
def all(self):
def all(self) -> FluentQuery:
"""返回 self 以支持链式调用(与 Django QuerySet.all() 一致)。
如需获取列表,使用 .to_list() 或 list()。"""
self._qs = self._qs.all()
return self
def to_list(self):
def to_list(self) -> List[models.Model]:
"""执行查询并返回模型实例列表。"""
if self._cache_enabled:
cache = _get_cache()
@@ -354,7 +400,8 @@ class FluentQuery:
_get_cache()[self._cache_key('to_list')] = results
return results
def count(self):
def count(self) -> int:
"""返回匹配记录数。"""
if self._cache_enabled:
cache = _get_cache()
key = self._cache_key('count')
@@ -365,7 +412,8 @@ class FluentQuery:
_get_cache()[self._cache_key('count')] = result
return result
def get(self, *args, **kwargs):
def get(self, *args: Any, **kwargs: Any) -> models.Model:
"""获取唯一匹配记录,不存在或多个时抛异常。"""
if args:
obj = self._qs.filter(*args).get(**kwargs)
else:
@@ -373,89 +421,108 @@ class FluentQuery:
_get_session().track(obj)
return obj
def create(self, **kwargs):
def create(self, **kwargs: Any) -> models.Model:
"""创建并保存新记录。"""
obj = self._model.objects.create(**kwargs)
_invalidate_cache()
return obj
def get_or_create(self, defaults=None, **kwargs):
def get_or_create(self, defaults: Optional[Dict[str, Any]] = None, **kwargs: Any) -> Tuple[models.Model, bool]:
"""查询或创建记录,返回 (obj, created) 元组。"""
obj, created = self._model.objects.get_or_create(defaults=defaults, **kwargs)
if created:
_invalidate_cache()
_get_session().track(obj)
return obj, created
def update_or_create(self, defaults=None, **kwargs):
def update_or_create(self, defaults: Optional[Dict[str, Any]] = None, **kwargs: Any) -> Tuple[models.Model, bool]:
"""更新或创建记录,返回 (obj, created) 元组。"""
obj, created = self._model.objects.update_or_create(defaults=defaults, **kwargs)
_invalidate_cache()
_get_session().track(obj)
return obj, created
def none(self):
def none(self) -> FluentQuery:
"""返回空结果集,返回 self 以支持链式调用。"""
self._qs = self._qs.none()
return self
def last(self):
def last(self) -> Optional[models.Model]:
"""返回最后一条记录或 None。"""
return self._qs.last()
def limit(self, n):
def limit(self, n: int) -> FluentQuery:
"""限制返回记录数,返回 self 以支持链式调用。"""
self._qs = self._qs[:n]
return self
def offset(self, n):
def offset(self, n: int) -> FluentQuery:
"""跳过前 n 条记录,返回 self 以支持链式调用。"""
self._qs = self._qs[n:]
return self
def join(self, *args):
def join(self, *args: Type[models.Model]) -> FluentQuery:
"""关联模型,用于跨表查询自动映射字段,返回 self 以支持链式调用。"""
for arg in args:
if isinstance(arg, type) and issubclass(arg, models.Model):
if arg not in self._joined_models:
self._joined_models.append(arg)
return self
def select_related(self, *fields):
def select_related(self, *fields: str) -> FluentQuery:
"""外键关联查询优化,返回 self 以支持链式调用。"""
self._qs = self._qs.select_related(*fields)
return self
def prefetch_related(self, *fields):
def prefetch_related(self, *fields: str) -> FluentQuery:
"""多对多/反向关联查询优化,返回 self 以支持链式调用。"""
self._qs = self._qs.prefetch_related(*fields)
return self
def only(self, *fields):
def only(self, *fields: str) -> FluentQuery:
"""只加载指定字段,返回 self 以支持链式调用。"""
self._qs = self._qs.only(*fields)
return self
def defer(self, *fields):
def defer(self, *fields: str) -> FluentQuery:
"""延迟加载指定字段,返回 self 以支持链式调用。"""
self._qs = self._qs.defer(*fields)
return self
def select_for_update(self, nowait=False, skip_locked=False, of=(), no_key=False):
def select_for_update(self, nowait: bool = False, skip_locked: bool = False,
of: Tuple = (), no_key: bool = False) -> FluentQuery:
"""行级锁,返回 self 以支持链式调用。"""
self._qs = self._qs.select_for_update(
nowait=nowait, skip_locked=skip_locked, of=of, no_key=no_key
)
return self
def bulk_create(self, objs, batch_size=None, ignore_conflicts=False):
def bulk_create(self, objs: List[models.Model], batch_size: Optional[int] = None,
ignore_conflicts: bool = False) -> List[models.Model]:
"""批量创建记录。"""
result = self._model.objects.bulk_create(
objs, batch_size=batch_size, ignore_conflicts=ignore_conflicts
)
_invalidate_cache()
return result
def annotate(self, **kwargs):
def annotate(self, **kwargs: Any) -> FluentQuery:
"""注解查询,添加聚合/计算字段,返回 self 以支持链式调用。"""
self._qs = self._qs.annotate(**kwargs)
return self
def aggregate(self, **kwargs):
def aggregate(self, **kwargs: Any) -> Dict[str, Any]:
"""聚合查询,返回聚合结果字典。"""
return self._qs.aggregate(**kwargs)
def values(self, *fields):
def values(self, *fields: Union[str, _AggregateExpr, models.Expression]) -> _ValuesResult:
"""返回指定字段的字典结果,支持链式调用。"""
if not fields:
# 无参数:返回 Django .values()(所有字段字典)
return _ValuesResult(self._qs.values(), [], [])
return _ValuesResult(self._qs.values(), [], fields)
annotations = {}
converted = []
annotations: Dict[str, models.Aggregate] = {}
converted: List[str] = []
for i, f in enumerate(fields):
if isinstance(f, _AggregateExpr):
annotations[f'val{i}'] = f.aggregate
@@ -475,33 +542,41 @@ class FluentQuery:
qs = qs.values(*converted)
return _ValuesResult(qs, converted, fields)
def with_entities(self, *fields):
def with_entities(self, *fields: Union[str, _AggregateExpr, models.Expression]) -> _ValuesResult:
"""values() 的别名SQLAlchemy 风格。"""
return self.values(*fields)
def params(self, **kwargs):
def params(self, **kwargs: Any) -> FluentQuery:
"""参数占位SQLAlchemy 兼容),返回 self 以支持链式调用。"""
return self
def values_list(self, *fields, flat=False):
def values_list(self, *fields: str, flat: bool = False) -> Any:
"""返回指定字段的元组列表。"""
return self._qs.values_list(*fields, flat=flat)
def update(self, **kwargs):
def update(self, **kwargs: Any) -> int:
"""批量更新匹配记录,返回受影响行数。"""
result = self._qs.update(**kwargs)
_invalidate_cache()
return result
def delete(self):
def delete(self) -> Tuple[int, Dict[str, int]]:
"""批量删除匹配记录,返回 (删除数, {表名: 删除数})。"""
result = self._qs.delete()
_invalidate_cache()
return result
def exists(self):
def exists(self) -> bool:
"""判断是否存在匹配记录。"""
return self._qs.exists()
def distinct(self):
def distinct(self) -> FluentQuery:
"""去重查询,返回 self 以支持链式调用。"""
self._qs = self._qs.distinct()
return self
def paginate(self, page=1, per_page=20, error_out=False):
def paginate(self, page: int = 1, per_page: int = 20, error_out: bool = False) -> Any:
"""分页查询,返回 Page 对象(含 items, total, pages 等属性)。"""
from django.core.paginator import Paginator
if not self._qs.ordered:
meta = getattr(self._model, '_meta', None)
@@ -523,16 +598,18 @@ class FluentQuery:
page_obj.prev_num = page_obj.previous_page_number() if page_obj.has_prev else None
return page_obj
def no_cache(self):
def no_cache(self) -> FluentQuery:
"""禁用查询缓存,返回 self 以支持链式调用。"""
self._cache_enabled = False
return self
def using(self, alias):
def using(self, alias: str) -> FluentQuery:
"""指定数据库别名,返回 self 以支持链式调用。"""
self._qs = self._qs.using(alias)
return self
def __getitem__(self, key):
"""支持切片 [start:end] 和索引 [n]"""
def __getitem__(self, key: Union[int, slice]) -> Any:
"""支持切片 [start:end] 和索引 [n]。切片返回 self索引返回模型实例。"""
if isinstance(key, slice):
self._qs = self._qs[key]
return self
@@ -540,16 +617,16 @@ class FluentQuery:
# 单个索引,返回模型实例
return self._qs[key]
def __iter__(self):
def __iter__(self) -> Iterator[models.Model]:
return iter(self._qs)
def __len__(self):
def __len__(self) -> int:
return self.count()
def __bool__(self):
def __bool__(self) -> bool:
return self.exists()
def __repr__(self):
def __repr__(self) -> str:
return f'FluentQuery({self._model.__name__})'

View File

@@ -1,3 +1,8 @@
from __future__ import annotations
import typing
from typing import Any, Dict, List, Optional, Tuple, Type, Union
from django.db import models
from django.db.models import Q
from .fluent import FluentQuery
@@ -290,9 +295,19 @@ class _CastJsonExpression:
class ForeignKeyIdDescriptor:
"""FK 字段描述符:类级别访问返回 FieldExpression实例级别访问/设置委托给原始 Django 字段。"""
def __init__(self, fk_name):
self.fk_name = fk_name
self.id_attr = fk_name + '_id'
self._original_field = None # 延迟绑定原始 Django 字段
def _get_original_field(self, owner):
if self._original_field is None:
try:
self._original_field = owner._meta.get_field(self.fk_name)
except Exception:
pass
return self._original_field
def __get__(self, obj, objtype=None):
if obj is None:
@@ -300,20 +315,39 @@ class ForeignKeyIdDescriptor:
return getattr(obj, self.id_attr)
def __set__(self, obj, value):
setattr(obj, self.id_attr, value)
# 直接写入实例 __dict__避免触发描述符递归
obj.__dict__[self.id_attr] = value
class FieldExpressionDescriptor:
"""普通字段描述符:类级别访问返回 FieldExpression实例级别访问/设置委托给原始 Django 字段。"""
def __init__(self, field_name):
self.field_name = field_name
self._original_field = None # 延迟绑定原始 Django 字段
def _get_original_field(self, owner):
if self._original_field is None:
try:
self._original_field = owner._meta.get_field(self.field_name)
except Exception:
pass
return self._original_field
def __get__(self, obj, objtype=None):
if obj is None:
return FieldExpression(objtype, self.field_name)
# 实例级别:从 __dict__ 获取,或让 Django 字段描述符处理
if self.field_name in obj.__dict__:
return obj.__dict__[self.field_name]
# 委托给原始 Django 字段
orig = self._get_original_field(objtype)
if orig is not None:
return orig.__get__(obj, objtype)
return getattr(obj, self.field_name)
def __set__(self, obj, value):
setattr(obj, self.field_name, value)
# 直接写入实例 __dict__避免触发描述符递归
obj.__dict__[self.field_name] = value
class QModelBase(models.base.ModelBase):
@@ -337,15 +371,90 @@ class QModelBase(models.base.ModelBase):
except Exception:
pass
return super().__new__(mcs, name, bases, namespace, **kwargs)
cls = super().__new__(mcs, name, bases, namespace, **kwargs)
# 模型创建后,自动为所有字段添加描述符,支持 Model.field_name 链式 OOP
if not cls._meta.abstract:
_attach_field_descriptors(cls)
return cls
def _attach_field_descriptors(cls: Type) -> None:
"""为 QModel 子类的所有字段自动添加 FieldExpressionDescriptor / ForeignKeyIdDescriptor
使 Model.field_name 在类级别访问时返回 FieldExpression支持链式 OOP 查询。"""
from django.db import models as dj_models
for field in cls._meta.local_fields:
fname = field.name
# 跳过已有描述符或内部属性
if fname.startswith('_'):
continue
# 检查是否已有同名描述符(避免覆盖手动定义的)
existing = cls.__dict__.get(fname)
if isinstance(existing, (FieldExpressionDescriptor, ForeignKeyIdDescriptor)):
continue
# 跳过 Django 自动创建的反向关系和 id 主键(如果用户未显式定义)
if fname == 'id' and field.primary_key and not isinstance(
cls.__dict__.get(fname), (FieldExpressionDescriptor, ForeignKeyIdDescriptor)
):
# id 字段也添加描述符
setattr(cls, fname, FieldExpressionDescriptor(fname))
continue
if isinstance(field, (dj_models.ForeignKey, dj_models.OneToOneField)):
# FK 字段:类级别访问返回 FieldExpression实例级别访问返回关联对象
setattr(cls, fname, ForeignKeyIdDescriptor(fname))
else:
setattr(cls, fname, FieldExpressionDescriptor(fname))
class QModel(models.Model, metaclass=QModelBase):
"""QModel 兼容层基类,提供 SQLAlchemy 风格的链式 OOP 查询接口。
所有继承 QModel 的模型自动获得:
- ``Model.query`` : 返回 FluentQuery支持链式查询
- ``Model.field_name`` : 类级别访问返回 FieldExpression支持表达式查询
- ``Model.field_name`` : 实例级别访问返回字段值(与普通 Django Model 一致)
用法示例::
# 链式查询
orders = Dingdan.query.filter(zhuangtai=8).order_by('-create_time').to_list()
# 字段表达式
Dingdan.query.filter(Dingdan.jine > 100, Dingdan.zhuangtai.in_([1,2,3]))
# 聚合
Dingdan.query.aggregate(total=Sum('jine'))
# values / annotate
Dingdan.query.values('zhuangtai').annotate(cnt=Count('id'))
# 分页
Dingdan.query.paginate(page=1, per_page=20)
"""
class Meta:
abstract = True
@classproperty
def query(cls):
def query(cls) -> FluentQuery:
"""返回 FluentQuery 兼容层查询构建器,支持链式 OOP 查询。"""
return FluentQuery(cls)
@classmethod
def __class_getitem__(cls, item) -> FluentQuery:
"""支持 Model[condition] 语法,返回 filter 后的 FluentQuery。
用法::
Dingdan[Dingdan.zhuangtai == 8].order_by('-create_time').to_list()
"""
if isinstance(item, Q):
return cls.query.filter(item)
if isinstance(item, dict):
return cls.query.filter(**item)
return cls.query.filter(item)