进行了 backend 的重构,实验中

This commit is contained in:
2026-06-18 04:15:23 +08:00
parent 54cc658f0b
commit 854c832c07
10 changed files with 1196 additions and 864 deletions

View File

@@ -1,149 +1,240 @@
"""FluentQuery 兼容层查询构建器
提供 SQLAlchemy 风格的链式 OOP 查询接口,包装 Django QuerySet。
核心类:
- ``FluentQuery`` — 链式查询构建器,由 ``Model.query`` 返回
- ``Session`` — 简易会话管理add/commit/rollback
- ``_CompatFunc`` — 兼容层聚合函数func.sum / func.count 等)
- ``db`` — 全局数据库操作入口db.session / db.or_ / db.and_
典型用法::
from gvsdsdk.fluent import db, func, FQ
# 链式查询
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)
# 兼容层聚合函数
func.sum(Dingdan.jine)
func.count(Dingdan.id)
func.year(Dingdan.create_time)
"""
from __future__ import annotations
import logging
import threading
import functools
from typing import Any, Dict, Iterator, List, Optional, Tuple, Type, Union, overload
from typing import Any, Dict, Generic, Iterator, List, Optional, Sequence, Tuple, Type, TypeVar, Union
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
logger = logging.getLogger(__name__)
_local = threading.local()
#: 模型类型变量,用于 FluentQuery[_M] 泛型参数
_M = TypeVar('_M', bound=models.Model)
def _get_field_expression_cls():
"""延迟导入,避免 fluent ↔ model_base 循环依赖"""
"""延迟导入 FieldExpression,避免 fluent ↔ model_base 循环依赖"""
from gvsdsdk.model_base import FieldExpression
return FieldExpression
# ---------------------------------------------------------------------------
# 兼容层聚合函数
# ---------------------------------------------------------------------------
class _AggregateExpr:
def __init__(self, model, aggregate):
"""聚合表达式包装,关联模型类和 Django Aggregate 对象。"""
def __init__(self, model: Optional[Type[models.Model]], aggregate: models.Aggregate) -> None:
self.model = model
self.aggregate = aggregate
class _ExtractExpr:
def __init__(self, field_name, lookup_type):
"""日期提取表达式year/month/day用于 ``func.year(field) == 2024`` 等。"""
def __init__(self, field_name: str, lookup_type: str) -> None:
self.field_name = field_name
self.lookup_type = lookup_type
def __eq__(self, other):
def __eq__(self, other) -> Q: # type: ignore[override]
return Q(**{f'{self.field_name}__{self.lookup_type}': other})
def __and__(self, other):
def __and__(self, other) -> Q:
if isinstance(other, Q):
return Q(**{f'{self.field_name}__{self.lookup_type}': None}) & other
return self.__eq__(other)
class _CompatFunc:
def sum(self, field_expr):
if hasattr(field_expr, '_django_name'):
agg = Sum(field_expr._django_name)
elif hasattr(field_expr, 'field_name'):
agg = Sum(field_expr.field_name)
else:
agg = Sum(str(field_expr))
model = getattr(field_expr, 'model_class', None) if hasattr(field_expr, 'model_class') else None
"""兼容层聚合函数集合,提供 SQLAlchemy 风格的 ``func.sum()`` / ``func.count()`` 等。
用法::
from gvsdsdk.fluent import func
func.sum(Model.jine) → Sum('jine')
func.count(Model.id) → Count('id')
func.max(Model.jine) → Max('jine')
func.min(Model.jine) → Min('jine')
func.avg(Model.jine) → Avg('jine')
func.year(Model.riqi) → _ExtractExpr('riqi', 'year')
func.month(Model.riqi) → _ExtractExpr('riqi', 'month')
func.day(Model.riqi) → _ExtractExpr('riqi', 'day')
"""
def sum(self, field_expr) -> Union[_AggregateExpr, Sum]:
"""求和聚合。"""
name = self._resolve_name(field_expr)
agg = Sum(name)
model = getattr(field_expr, 'model_class', None)
return _AggregateExpr(model, agg) if model else agg
def count(self, field_expr=None):
def count(self, field_expr=None) -> Union[_AggregateExpr, Count]:
"""计数聚合,无参数时等价于 COUNT(*)。"""
if field_expr is None:
return Count('*')
if hasattr(field_expr, '_django_name'):
agg = Count(field_expr._django_name)
else:
agg = Count(str(field_expr))
model = getattr(field_expr, 'model_class', None) if hasattr(field_expr, 'model_class') else None
name = self._resolve_name(field_expr)
agg = Count(name)
model = getattr(field_expr, 'model_class', None)
return _AggregateExpr(model, agg) if model else agg
def max(self, field_expr):
if hasattr(field_expr, '_django_name'):
agg = Max(field_expr._django_name)
else:
agg = Max(str(field_expr))
model = getattr(field_expr, 'model_class', None) if hasattr(field_expr, 'model_class') else None
def max(self, field_expr) -> Union[_AggregateExpr, Max]:
"""最大值聚合。"""
name = self._resolve_name(field_expr)
agg = Max(name)
model = getattr(field_expr, 'model_class', None)
return _AggregateExpr(model, agg) if model else agg
def min(self, field_expr):
if hasattr(field_expr, '_django_name'):
agg = Min(field_expr._django_name)
else:
agg = Min(str(field_expr))
model = getattr(field_expr, 'model_class', None) if hasattr(field_expr, 'model_class') else None
def min(self, field_expr) -> Union[_AggregateExpr, Min]:
"""最小值聚合。"""
name = self._resolve_name(field_expr)
agg = Min(name)
model = getattr(field_expr, 'model_class', None)
return _AggregateExpr(model, agg) if model else agg
def avg(self, field_expr):
if hasattr(field_expr, '_django_name'):
agg = Avg(field_expr._django_name)
else:
agg = Avg(str(field_expr))
model = getattr(field_expr, 'model_class', None) if hasattr(field_expr, 'model_class') else None
def avg(self, field_expr) -> Union[_AggregateExpr, Avg]:
"""平均值聚合。"""
name = self._resolve_name(field_expr)
agg = Avg(name)
model = getattr(field_expr, 'model_class', None)
return _AggregateExpr(model, agg) if model else agg
def year(self, field_expr):
if hasattr(field_expr, '_django_name'):
return _ExtractExpr(field_expr._django_name, 'year')
return _ExtractExpr(str(field_expr), 'year')
def year(self, field_expr) -> _ExtractExpr:
"""提取年份,用于 ``func.year(Model.date) == 2024``。"""
return _ExtractExpr(self._resolve_name(field_expr), 'year')
def month(self, field_expr):
if hasattr(field_expr, '_django_name'):
return _ExtractExpr(field_expr._django_name, 'month')
return _ExtractExpr(str(field_expr), 'month')
def month(self, field_expr) -> _ExtractExpr:
"""提取月份。"""
return _ExtractExpr(self._resolve_name(field_expr), 'month')
def day(self, field_expr):
def day(self, field_expr) -> _ExtractExpr:
"""提取日期。"""
return _ExtractExpr(self._resolve_name(field_expr), 'day')
@staticmethod
def _resolve_name(field_expr) -> str:
"""从 FieldExpression 或字符串解析 Django 字段名。"""
if hasattr(field_expr, '_django_name'):
return _ExtractExpr(field_expr._django_name, 'day')
return _ExtractExpr(str(field_expr), 'day')
return field_expr._django_name
if hasattr(field_expr, 'field_name'):
return field_expr.field_name
return str(field_expr)
func = _CompatFunc()
# ---------------------------------------------------------------------------
# DELETE 语句构建器
# ---------------------------------------------------------------------------
class _DeleteStmt:
def __init__(self, model_class, conditions=None):
"""DELETE 语句构建器,支持 ``Model.query.filter(...).delete()`` 风格。"""
def __init__(self, model_class: Type[models.Model], conditions: Optional[Q] = None) -> None:
self.model_class = model_class
self.conditions = conditions or Q()
def where(self, *args, **kwargs):
def where(self, *args: Q, **kwargs: Any) -> _DeleteStmt:
"""添加 WHERE 条件。"""
new_conditions = Q()
for a in args:
if isinstance(a, Q):
new_conditions &= a
if kwargs:
new_conditions &= Q(**kwargs)
self.conditions = new_conditions
return self
def execute(self):
def execute(self) -> Tuple[int, Dict[str, int]]:
"""执行删除,返回 (删除数, {表名: 删除数})。"""
if self.model_class and self.conditions:
self.model_class.objects.filter(self.conditions).delete()
return self.model_class.objects.filter(self.conditions).delete()
elif self.model_class:
self.model_class.objects.all().delete()
return self.model_class.objects.all().delete()
return (0, {})
def _get_session():
# ---------------------------------------------------------------------------
# 会话与缓存
# ---------------------------------------------------------------------------
def _get_session() -> Session:
"""获取当前线程的 Session 实例。"""
if not hasattr(_local, 'db_session'):
_local.db_session = Session()
return _local.db_session
def _get_cache():
def _get_cache() -> Dict:
"""获取当前线程的查询缓存。"""
if not hasattr(_local, 'query_cache'):
_local.query_cache = {}
return _local.query_cache
def _invalidate_cache():
def _invalidate_cache() -> None:
"""清除当前线程的查询缓存。"""
if hasattr(_local, 'query_cache'):
_local.query_cache.clear()
# ---------------------------------------------------------------------------
# 关联字段访问器
# ---------------------------------------------------------------------------
class RelationAccessor:
def __init__(self, model_class, prefix=''):
"""关联字段访问器,支持 ``Model.relation.field`` 跨表字段表达式。
用法::
# 自动解析关联字段
Dingdan.query.filter(Dingdan.shangjia.nicheng == '测试')
"""
def __init__(self, model_class: Type[models.Model], prefix: str = '') -> None:
self._model = model_class
self._prefix = prefix
def __getattr__(self, name):
def __getattr__(self, name: str):
if name.startswith('_'):
return object.__getattribute__(self, name)
try:
@@ -163,18 +254,26 @@ class RelationAccessor:
return RelationAccessor(rel.related_model, full_prefix)
return _get_field_expression_cls()(self._model, name)
def __repr__(self) -> str:
return f'RelationAccessor({self._model.__name__}, prefix={self._prefix!r})'
# ---------------------------------------------------------------------------
# Values 结果集
# ---------------------------------------------------------------------------
class _ValuesResult:
"""包装 Django ValuesQuerySet支持链式调用。
底层使用 Django ``.values()`` 返回字典,保持与 Django 行为一致。
额外提供 ``.all()`` 返回字典列表,``.first()`` 返回字典或 None
``.scalar()`` 返回单字段场景的标量值。
"""
包装 Django ValuesQuerySet支持链式调用。
底层使用 Django .values() 返回字典,保持与 Django 行为一致。
额外提供 .all() 返回字典列表,.first() 返回字典或 None
.scalar() 返回单字段场景的标量值。
"""
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 # 原始传入的字段(用于判断单字段场景)
self._qs: models.QuerySet = qs
self._fields: List[str] = fields
self._original_fields: Tuple = original_fields
@property
def query(self) -> Any:
@@ -202,12 +301,15 @@ class _ValuesResult:
return self
def first(self) -> Optional[Dict[str, Any]]:
"""返回第一条记录(字典)或 None。"""
return self._qs.first()
def all(self) -> List[Dict[str, Any]]:
"""返回所有记录的字典列表。"""
return list(self._qs)
def scalar(self) -> Any:
"""返回单字段场景的标量值。"""
row = self._qs.first()
if row is None:
return None
@@ -237,11 +339,18 @@ class _ValuesResult:
def __bool__(self) -> bool:
return self._qs.exists()
def __repr__(self) -> str:
return f'_ValuesResult(fields={self._fields}, count={self.count()})'
class FluentQuery:
# ---------------------------------------------------------------------------
# FluentQuery 核心查询构建器
# ---------------------------------------------------------------------------
class FluentQuery(Generic[_M]):
"""QModel 兼容层查询构建器,提供 SQLAlchemy 风格的链式 OOP 查询接口。
法示例::
由 ``Model.query`` 返回,支持链式调用::
# 基础查询
Dingdan.query.filter(zhuangtai=8).all()
@@ -249,7 +358,7 @@ class FluentQuery:
# 链式过滤
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]))
# 聚合
@@ -260,10 +369,12 @@ class FluentQuery:
# 分页
Dingdan.query.paginate(page=1, per_page=20)
泛型参数 ``_M`` 携带模型类型,使 IDE 能推断链式调用返回的模型实例类型。
"""
def __init__(self, model_class: Type[models.Model]) -> None:
self._model: Type[models.Model] = model_class
def __init__(self, model_class: Type[_M]) -> None:
self._model: Type[_M] = model_class
self._qs: models.QuerySet = model_class.objects.all()
self._joined_models: List[Type[models.Model]] = []
self._cache_enabled: bool = True
@@ -273,7 +384,13 @@ class FluentQuery:
"""暴露底层 Django QuerySet 的 query 属性,支持 Subquery() 等场景。"""
return self._qs.query
def _clone(self) -> FluentQuery:
@property
def model(self) -> Type[_M]:
"""返回关联的模型类。"""
return self._model
def _clone(self) -> FluentQuery[_M]:
"""克隆当前查询状态。"""
q = FluentQuery(self._model)
q._qs = self._qs
q._joined_models = list(self._joined_models)
@@ -281,6 +398,7 @@ class FluentQuery:
return q
def _cache_key(self, method: str) -> Tuple[Any, str, Any]:
"""生成缓存键。"""
try:
query_str = str(self._qs.query)
except Exception:
@@ -288,6 +406,7 @@ class FluentQuery:
return (self._model, method, query_str)
def _remap_q(self, q: Q) -> Q:
"""重映射 Q 对象中的字段名(用于 join 跨表查询)。"""
new_children: List[Any] = []
for child in q.children:
if isinstance(child, Q):
@@ -306,6 +425,7 @@ class FluentQuery:
return q
def _remap_key(self, key: str) -> str:
"""重映射字段名(用于 join 跨表查询自动添加关联前缀)。"""
parts = key.split('__')
field_name = parts[0]
try:
@@ -330,7 +450,9 @@ class FluentQuery:
pass
return key
def filter(self, *args: Any, **kwargs: Any) -> FluentQuery:
# ---- 过滤 ----
def filter(self, *args: Any, **kwargs: Any) -> FluentQuery[_M]:
"""过滤查询,支持 Q 对象和关键字参数,返回 self 以支持链式调用。"""
converted_args: List[Any] = []
for arg in args:
@@ -341,7 +463,7 @@ class FluentQuery:
self._qs = self._qs.filter(*converted_args, **kwargs)
return self
def exclude(self, *args: Any, **kwargs: Any) -> FluentQuery:
def exclude(self, *args: Any, **kwargs: Any) -> FluentQuery[_M]:
"""排除查询,返回 self 以支持链式调用。"""
converted_args: List[Any] = []
for arg in args:
@@ -352,7 +474,9 @@ class FluentQuery:
self._qs = self._qs.exclude(*converted_args, **kwargs)
return self
def order_by(self, *args: Any) -> FluentQuery:
# ---- 排序 ----
def order_by(self, *args: Any) -> FluentQuery[_M]:
"""排序,支持字符串和 FieldExpression返回 self 以支持链式调用。"""
converted: List[str] = []
for a in args:
@@ -365,7 +489,9 @@ class FluentQuery:
self._qs = self._qs.order_by(*converted)
return self
def first(self) -> Optional[models.Model]:
# ---- 获取记录 ----
def first(self) -> Optional[_M]:
"""返回第一条记录或 None。"""
if self._cache_enabled:
cache = _get_cache()
@@ -379,13 +505,26 @@ class FluentQuery:
_get_cache()[self._cache_key('first')] = obj
return obj
def all(self) -> FluentQuery:
def last(self) -> Optional[_M]:
"""返回最后一条记录或 None。"""
return self._qs.last()
def get(self, *args: Any, **kwargs: Any) -> _M:
"""获取唯一匹配记录,不存在或多个时抛异常。"""
if args:
obj = self._qs.filter(*args).get(**kwargs)
else:
obj = self._qs.get(**kwargs)
_get_session().track(obj)
return obj
def all(self) -> FluentQuery[_M]:
"""返回 self 以支持链式调用(与 Django QuerySet.all() 一致)。
如需获取列表,使用 .to_list() 或 list()。"""
如需获取列表,使用 ``.to_list()````list()``"""
self._qs = self._qs.all()
return self
def to_list(self) -> List[models.Model]:
def to_list(self) -> List[_M]:
"""执行查询并返回模型实例列表。"""
if self._cache_enabled:
cache = _get_cache()
@@ -412,93 +551,37 @@ class FluentQuery:
_get_cache()[self._cache_key('count')] = result
return result
def get(self, *args: Any, **kwargs: Any) -> models.Model:
"""获取唯一匹配记录,不存在或多个时抛异常"""
if args:
obj = self._qs.filter(*args).get(**kwargs)
else:
obj = self._qs.get(**kwargs)
_get_session().track(obj)
return obj
def exists(self) -> bool:
"""判断是否存在匹配记录"""
return self._qs.exists()
def create(self, **kwargs: Any) -> models.Model:
# ---- 创建/更新/删除 ----
def create(self, **kwargs: Any) -> _M:
"""创建并保存新记录。"""
obj = self._model.objects.create(**kwargs)
_invalidate_cache()
return obj
def get_or_create(self, defaults: Optional[Dict[str, Any]] = None, **kwargs: Any) -> Tuple[models.Model, bool]:
"""查询或创建记录,返回 (obj, created) 元组。"""
def get_or_create(self, defaults: Optional[Dict[str, Any]] = None,
**kwargs: Any) -> Tuple[_M, 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: Optional[Dict[str, Any]] = None, **kwargs: Any) -> Tuple[models.Model, bool]:
"""更新或创建记录,返回 (obj, created) 元组。"""
def update_or_create(self, defaults: Optional[Dict[str, Any]] = None,
**kwargs: Any) -> Tuple[_M, 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) -> FluentQuery:
"""返回空结果集,返回 self 以支持链式调用。"""
self._qs = self._qs.none()
return self
def last(self) -> Optional[models.Model]:
"""返回最后一条记录或 None。"""
return self._qs.last()
def limit(self, n: int) -> FluentQuery:
"""限制返回记录数,返回 self 以支持链式调用。"""
self._qs = self._qs[:n]
return self
def offset(self, n: int) -> FluentQuery:
"""跳过前 n 条记录,返回 self 以支持链式调用。"""
self._qs = self._qs[n:]
return self
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: str) -> FluentQuery:
"""外键关联查询优化,返回 self 以支持链式调用。"""
self._qs = self._qs.select_related(*fields)
return self
def prefetch_related(self, *fields: str) -> FluentQuery:
"""多对多/反向关联查询优化,返回 self 以支持链式调用。"""
self._qs = self._qs.prefetch_related(*fields)
return self
def only(self, *fields: str) -> FluentQuery:
"""只加载指定字段,返回 self 以支持链式调用。"""
self._qs = self._qs.only(*fields)
return self
def defer(self, *fields: str) -> FluentQuery:
"""延迟加载指定字段,返回 self 以支持链式调用。"""
self._qs = self._qs.defer(*fields)
return self
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: List[models.Model], batch_size: Optional[int] = None,
ignore_conflicts: bool = False) -> List[models.Model]:
def bulk_create(self, objs: List[_M], batch_size: Optional[int] = None,
ignore_conflicts: bool = False) -> List[_M]:
"""批量创建记录。"""
result = self._model.objects.bulk_create(
objs, batch_size=batch_size, ignore_conflicts=ignore_conflicts
@@ -506,7 +589,100 @@ class FluentQuery:
_invalidate_cache()
return result
def annotate(self, **kwargs: Any) -> FluentQuery:
def bulk_update(self, objs: List[_M], fields: List[str],
batch_size: Optional[int] = None) -> int:
"""批量更新记录,返回受影响行数。"""
result = self._model.objects.bulk_update(objs, fields, batch_size=batch_size)
_invalidate_cache()
return result
def update(self, **kwargs: Any) -> int:
"""批量更新匹配记录,返回受影响行数。"""
result = self._qs.update(**kwargs)
_invalidate_cache()
return result
def delete(self) -> Tuple[int, Dict[str, int]]:
"""批量删除匹配记录,返回 ``(删除数, {表名: 删除数})``。"""
result = self._qs.delete()
_invalidate_cache()
return result
# ---- 查询修饰 ----
def none(self) -> FluentQuery[_M]:
"""返回空结果集,返回 self 以支持链式调用。"""
self._qs = self._qs.none()
return self
def limit(self, n: int) -> FluentQuery[_M]:
"""限制返回记录数,返回 self 以支持链式调用。"""
self._qs = self._qs[:n]
return self
def offset(self, n: int) -> FluentQuery[_M]:
"""跳过前 n 条记录,返回 self 以支持链式调用。"""
self._qs = self._qs[n:]
return self
def join(self, *args: Type[models.Model]) -> FluentQuery[_M]:
"""关联模型,用于跨表查询自动映射字段,返回 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: str) -> FluentQuery[_M]:
"""外键关联查询优化,返回 self 以支持链式调用。"""
self._qs = self._qs.select_related(*fields)
return self
def prefetch_related(self, *fields: str) -> FluentQuery[_M]:
"""多对多/反向关联查询优化,返回 self 以支持链式调用。"""
self._qs = self._qs.prefetch_related(*fields)
return self
def only(self, *fields: str) -> FluentQuery[_M]:
"""只加载指定字段,返回 self 以支持链式调用。"""
self._qs = self._qs.only(*fields)
return self
def defer(self, *fields: str) -> FluentQuery[_M]:
"""延迟加载指定字段,返回 self 以支持链式调用。"""
self._qs = self._qs.defer(*fields)
return self
def select_for_update(self, nowait: bool = False, skip_locked: bool = False,
of: Tuple = (), no_key: bool = False) -> FluentQuery[_M]:
"""行级锁,返回 self 以支持链式调用。"""
self._qs = self._qs.select_for_update(
nowait=nowait, skip_locked=skip_locked, of=of, no_key=no_key
)
return self
def distinct(self) -> FluentQuery[_M]:
"""去重查询,返回 self 以支持链式调用。"""
self._qs = self._qs.distinct()
return self
def using(self, alias: str) -> FluentQuery[_M]:
"""指定数据库别名,返回 self 以支持链式调用。"""
self._qs = self._qs.using(alias)
return self
def no_cache(self) -> FluentQuery[_M]:
"""禁用查询缓存,返回 self 以支持链式调用。"""
self._cache_enabled = False
return self
def params(self, **kwargs: Any) -> FluentQuery[_M]:
"""参数占位SQLAlchemy 兼容),返回 self 以支持链式调用。"""
return self
# ---- 聚合与注解 ----
def annotate(self, **kwargs: Any) -> FluentQuery[_M]:
"""注解查询,添加聚合/计算字段,返回 self 以支持链式调用。"""
self._qs = self._qs.annotate(**kwargs)
return self
@@ -518,7 +694,6 @@ class FluentQuery:
def values(self, *fields: Union[str, _AggregateExpr, models.Expression]) -> _ValuesResult:
"""返回指定字段的字典结果,支持链式调用。"""
if not fields:
# 无参数:返回 Django .values()(所有字段字典)
return _ValuesResult(self._qs.values(), [], fields)
annotations: Dict[str, models.Aggregate] = {}
@@ -538,7 +713,6 @@ class FluentQuery:
qs = self._qs
if annotations:
qs = qs.annotate(**annotations)
# 使用 Django 原生 .values() 返回字典
qs = qs.values(*converted)
return _ValuesResult(qs, converted, fields)
@@ -546,38 +720,16 @@ class FluentQuery:
"""values() 的别名SQLAlchemy 风格。"""
return self.values(*fields)
def params(self, **kwargs: Any) -> FluentQuery:
"""参数占位SQLAlchemy 兼容),返回 self 以支持链式调用。"""
return self
def values_list(self, *fields: str, flat: bool = False) -> Any:
"""返回指定字段的元组列表。"""
return self._qs.values_list(*fields, flat=flat)
def update(self, **kwargs: Any) -> int:
"""批量更新匹配记录,返回受影响行数。"""
result = self._qs.update(**kwargs)
_invalidate_cache()
return result
# ---- 分页 ----
def delete(self) -> Tuple[int, Dict[str, int]]:
"""批量删除匹配记录,返回 (删除数, {表名: 删除数})。"""
result = self._qs.delete()
_invalidate_cache()
return result
def exists(self) -> bool:
"""判断是否存在匹配记录。"""
return self._qs.exists()
def distinct(self) -> FluentQuery:
"""去重查询,返回 self 以支持链式调用。"""
self._qs = self._qs.distinct()
return self
def paginate(self, page: int = 1, per_page: int = 20, error_out: bool = False) -> Any:
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
from django.core.paginator import Paginator, EmptyPage
if not self._qs.ordered:
meta = getattr(self._model, '_meta', None)
if meta and meta.pk:
@@ -585,7 +737,7 @@ class FluentQuery:
paginator = Paginator(self._qs, per_page)
try:
page_obj = paginator.page(page)
except Exception:
except EmptyPage:
page_obj = paginator.page(1)
page_obj.total = paginator.count
page_obj.items = page_obj.object_list
@@ -598,26 +750,17 @@ 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) -> FluentQuery:
"""禁用查询缓存,返回 self 以支持链式调用。"""
self._cache_enabled = False
return self
def using(self, alias: str) -> FluentQuery:
"""指定数据库别名,返回 self 以支持链式调用。"""
self._qs = self._qs.using(alias)
return self
# ---- 魔术方法 ----
def __getitem__(self, key: Union[int, slice]) -> Any:
"""支持切片 [start:end] 和索引 [n]。切片返回 self索引返回模型实例。"""
"""支持切片 ``[start:end]`` 和索引 ``[n]``。切片返回 self索引返回模型实例。"""
if isinstance(key, slice):
self._qs = self._qs[key]
return self
else:
# 单个索引,返回模型实例
return self._qs[key]
def __iter__(self) -> Iterator[models.Model]:
def __iter__(self) -> Iterator[_M]:
return iter(self._qs)
def __len__(self) -> int:
@@ -630,25 +773,43 @@ class FluentQuery:
return f'FluentQuery({self._model.__name__})'
class Session:
def __init__(self):
self._pending_adds = []
self._pending_deletes = []
self._tracked = set()
# ---------------------------------------------------------------------------
# Session 会话管理
# ---------------------------------------------------------------------------
def add(self, obj):
class Session:
"""简易会话管理器,支持 add/delete/commit/rollback 操作。
用法::
session = db.session
session.add(obj)
session.commit()
"""
def __init__(self) -> None:
self._pending_adds: List[models.Model] = []
self._pending_deletes: List[models.Model] = []
self._tracked: set = set()
def add(self, obj: models.Model) -> None:
"""添加对象到待保存列表。"""
self._pending_adds.append(obj)
def delete(self, obj):
def delete(self, obj: models.Model) -> None:
"""添加对象到待删除列表。"""
self._pending_deletes.append(obj)
self._tracked.discard(obj)
def track(self, obj):
def track(self, obj: models.Model) -> None:
"""跟踪对象变更commit 时自动 save"""
if isinstance(obj, models.Model):
self._tracked.add(obj)
def commit(self):
delete_pks = {(obj.__class__, obj.pk) for obj in self._pending_deletes if obj.pk is not None}
def commit(self) -> None:
"""提交所有待处理操作(删除 → 新增 → 更新)。"""
delete_pks = {(obj.__class__, obj.pk) for obj in self._pending_deletes
if obj.pk is not None}
with transaction.atomic():
for obj in self._pending_deletes:
obj.delete()
@@ -660,21 +821,23 @@ class Session:
try:
obj.save()
except Exception:
pass
logger.debug('Session.commit: skip tracked obj %s', obj, exc_info=True)
self._pending_adds.clear()
self._pending_deletes.clear()
self._tracked.clear()
_invalidate_cache()
def rollback(self):
def rollback(self) -> None:
"""回滚所有待处理操作。"""
self._pending_adds.clear()
self._pending_deletes.clear()
self._tracked.clear()
_invalidate_cache()
def query(self, *args):
model = None
aggregates = []
def query(self, *args) -> Optional[FluentQuery[Any]]:
"""创建查询构建器,支持传入模型类或聚合表达式。"""
model: Optional[Type[models.Model]] = None
aggregates: List = []
for a in args:
if isinstance(a, type) and issubclass(a, models.Model):
model = a
@@ -691,10 +854,12 @@ class Session:
return fq.values(*args)
return None
def flush(self):
def flush(self) -> None:
"""flush() 的别名,等同于 commit()。"""
self.commit()
def execute(self, stmt):
def execute(self, stmt) -> None:
"""执行语句(支持 _DeleteStmt 和原始 SQL 字符串)。"""
if isinstance(stmt, _DeleteStmt):
stmt.execute()
elif isinstance(stmt, str):
@@ -703,14 +868,31 @@ class Session:
cursor.execute(stmt)
# ---------------------------------------------------------------------------
# 全局入口
# ---------------------------------------------------------------------------
class DB:
"""全局数据库操作入口。
用法::
from gvsdsdk.fluent import db
session = db.session
q = db.or_(Model.field == 1, Model.field == 2)
q = db.and_(Model.field > 0, Model.field < 100)
"""
@property
def session(self):
def session(self) -> Session:
"""获取当前线程的 Session 实例。"""
return _get_session()
@staticmethod
def or_(*args):
result = None
def or_(*args: Q) -> Q:
"""逻辑 OR合并多个 Q 对象。"""
result: Optional[Q] = None
for q in args:
if isinstance(q, Q):
if result is None:
@@ -720,7 +902,8 @@ class DB:
return result if result is not None else Q()
@staticmethod
def and_(*args):
def and_(*args: Q) -> Q:
"""逻辑 AND合并多个 Q 对象。"""
result = Q()
for q in args:
if isinstance(q, Q):
@@ -731,5 +914,6 @@ class DB:
db = DB()
def FQ(model_class):
def FQ(model_class: Type[_M]) -> FluentQuery[_M]:
"""FluentQuery 工厂函数,等同于 ``Model.query``。"""
return FluentQuery(model_class)