修复了 QModel 链式表存在的一些问题
This commit is contained in:
@@ -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__})'
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user