From 54a64dc220fc9a3e87bd0d57ce756c0ad7c25a87 Mon Sep 17 00:00:00 2001 From: TermiNexus Date: Tue, 16 Jun 2026 23:33:23 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E4=BA=86=20QModel=20?= =?UTF-8?q?=E9=93=BE=E5=BC=8F=E8=A1=A8=E5=AD=98=E5=9C=A8=E7=9A=84=E4=B8=80?= =?UTF-8?q?=E4=BA=9B=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- gvsdsdk/fluent.py | 235 ++++++++++++++++++++++++++++-------------- gvsdsdk/model_base.py | 117 ++++++++++++++++++++- 2 files changed, 269 insertions(+), 83 deletions(-) diff --git a/gvsdsdk/fluent.py b/gvsdsdk/fluent.py index 6306adf..f6d007e 100644 --- a/gvsdsdk/fluent.py +++ b/gvsdsdk/fluent.py @@ -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__})' diff --git a/gvsdsdk/model_base.py b/gvsdsdk/model_base.py index 2a17f9a..34b2d9f 100644 --- a/gvsdsdk/model_base.py +++ b/gvsdsdk/model_base.py @@ -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) +