加入了 GVSDSDK 模块,进行了 QModel 兼容层的尝试,生产环境可用
This commit is contained in:
651
gvsdsdk/fluent.py
Normal file
651
gvsdsdk/fluent.py
Normal file
@@ -0,0 +1,651 @@
|
||||
import threading
|
||||
import functools
|
||||
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()
|
||||
|
||||
|
||||
class _AggregateExpr:
|
||||
def __init__(self, model, aggregate):
|
||||
self.model = model
|
||||
self.aggregate = aggregate
|
||||
|
||||
|
||||
class _ExtractExpr:
|
||||
def __init__(self, field_name, lookup_type):
|
||||
self.field_name = field_name
|
||||
self.lookup_type = lookup_type
|
||||
|
||||
def __eq__(self, other):
|
||||
return Q(**{f'{self.field_name}__{self.lookup_type}': other})
|
||||
|
||||
def __and__(self, other):
|
||||
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
|
||||
return _AggregateExpr(model, agg) if model else agg
|
||||
|
||||
def count(self, field_expr=None):
|
||||
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
|
||||
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
|
||||
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
|
||||
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
|
||||
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 month(self, field_expr):
|
||||
if hasattr(field_expr, '_django_name'):
|
||||
return _ExtractExpr(field_expr._django_name, 'month')
|
||||
return _ExtractExpr(str(field_expr), 'month')
|
||||
|
||||
def day(self, field_expr):
|
||||
if hasattr(field_expr, '_django_name'):
|
||||
return _ExtractExpr(field_expr._django_name, 'day')
|
||||
return _ExtractExpr(str(field_expr), 'day')
|
||||
|
||||
|
||||
func = _CompatFunc()
|
||||
|
||||
|
||||
class _DeleteStmt:
|
||||
def __init__(self, model_class, conditions=None):
|
||||
self.model_class = model_class
|
||||
self.conditions = conditions or Q()
|
||||
|
||||
def where(self, *args, **kwargs):
|
||||
new_conditions = Q()
|
||||
for a in args:
|
||||
if isinstance(a, Q):
|
||||
new_conditions &= a
|
||||
self.conditions = new_conditions
|
||||
return self
|
||||
|
||||
def execute(self):
|
||||
if self.model_class and self.conditions:
|
||||
self.model_class.objects.filter(self.conditions).delete()
|
||||
elif self.model_class:
|
||||
self.model_class.objects.all().delete()
|
||||
|
||||
|
||||
def _get_session():
|
||||
if not hasattr(_local, 'db_session'):
|
||||
_local.db_session = Session()
|
||||
return _local.db_session
|
||||
|
||||
|
||||
def _get_cache():
|
||||
if not hasattr(_local, 'query_cache'):
|
||||
_local.query_cache = {}
|
||||
return _local.query_cache
|
||||
|
||||
|
||||
def _invalidate_cache():
|
||||
if hasattr(_local, 'query_cache'):
|
||||
_local.query_cache.clear()
|
||||
|
||||
|
||||
class RelationAccessor:
|
||||
def __init__(self, model_class, prefix=''):
|
||||
self._model = model_class
|
||||
self._prefix = prefix
|
||||
|
||||
def __getattr__(self, name):
|
||||
if name.startswith('_'):
|
||||
return object.__getattribute__(self, name)
|
||||
try:
|
||||
self._model._meta.get_field(name)
|
||||
full_name = f'{self._prefix}__{name}' if self._prefix else name
|
||||
return FieldExpression(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:
|
||||
rel_name = ''
|
||||
if hasattr(rel, 'get_accessor_name'):
|
||||
rel_name = rel.get_accessor_name()
|
||||
elif hasattr(rel, 'name'):
|
||||
rel_name = rel.name
|
||||
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)
|
||||
|
||||
|
||||
class _ValuesResult:
|
||||
"""
|
||||
包装 Django ValuesQuerySet,支持链式调用。
|
||||
底层使用 Django .values() 返回字典,保持与 Django 行为一致。
|
||||
额外提供 .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 filter(self, *args, **kwargs):
|
||||
self._qs = self._qs.filter(*args, **kwargs)
|
||||
return self
|
||||
|
||||
def exclude(self, *args, **kwargs):
|
||||
self._qs = self._qs.exclude(*args, **kwargs)
|
||||
return self
|
||||
|
||||
def order_by(self, *fields):
|
||||
self._qs = self._qs.order_by(*fields)
|
||||
return self
|
||||
|
||||
def distinct(self):
|
||||
self._qs = self._qs.distinct()
|
||||
return self
|
||||
|
||||
def annotate(self, **kwargs):
|
||||
self._qs = self._qs.annotate(**kwargs)
|
||||
return self
|
||||
|
||||
def first(self):
|
||||
return self._qs.first()
|
||||
|
||||
def all(self):
|
||||
return list(self._qs)
|
||||
|
||||
def scalar(self):
|
||||
row = self._qs.first()
|
||||
if row is None:
|
||||
return None
|
||||
if self._fields:
|
||||
return row.get(self._fields[0])
|
||||
return row
|
||||
|
||||
def count(self):
|
||||
return self._qs.count()
|
||||
|
||||
def exists(self):
|
||||
return self._qs.exists()
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self._qs)
|
||||
|
||||
def __len__(self):
|
||||
return self.count()
|
||||
|
||||
def __bool__(self):
|
||||
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
|
||||
|
||||
def _clone(self):
|
||||
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):
|
||||
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 = []
|
||||
for child in q.children:
|
||||
if isinstance(child, Q):
|
||||
new_children.append(self._remap_q(child))
|
||||
elif isinstance(child, tuple):
|
||||
key, val = child
|
||||
new_key = self._remap_key(key)
|
||||
new_children.append((new_key, val))
|
||||
else:
|
||||
new_children.append(child)
|
||||
if new_children:
|
||||
new_q = Q(*new_children)
|
||||
new_q.connector = q.connector
|
||||
new_q.negated = q.negated
|
||||
return new_q
|
||||
return q
|
||||
|
||||
def _remap_key(self, key):
|
||||
parts = key.split('__')
|
||||
field_name = parts[0]
|
||||
try:
|
||||
self._model._meta.get_field(field_name)
|
||||
return key
|
||||
except Exception:
|
||||
pass
|
||||
for joined in self._joined_models:
|
||||
try:
|
||||
joined._meta.get_field(field_name)
|
||||
for rel in self._model._meta.get_fields(include_hidden=True):
|
||||
if hasattr(rel, 'related_model') and rel.related_model == joined:
|
||||
acc = ''
|
||||
if hasattr(rel, 'get_accessor_name'):
|
||||
acc = rel.get_accessor_name()
|
||||
elif hasattr(rel, 'name'):
|
||||
acc = rel.name
|
||||
if acc and acc != '+':
|
||||
return acc + '__' + key
|
||||
return joined._meta.model_name + '__' + key
|
||||
except Exception:
|
||||
pass
|
||||
return key
|
||||
|
||||
def filter(self, *args, **kwargs):
|
||||
converted_args = []
|
||||
for arg in args:
|
||||
if isinstance(arg, Q) and self._joined_models:
|
||||
converted_args.append(self._remap_q(arg))
|
||||
else:
|
||||
converted_args.append(arg)
|
||||
self._qs = self._qs.filter(*converted_args, **kwargs)
|
||||
return self
|
||||
|
||||
def exclude(self, *args, **kwargs):
|
||||
converted_args = []
|
||||
for arg in args:
|
||||
if isinstance(arg, Q) and self._joined_models:
|
||||
converted_args.append(self._remap_q(arg))
|
||||
else:
|
||||
converted_args.append(arg)
|
||||
self._qs = self._qs.exclude(*converted_args, **kwargs)
|
||||
return self
|
||||
|
||||
def order_by(self, *args):
|
||||
converted = []
|
||||
for a in args:
|
||||
if isinstance(a, str):
|
||||
converted.append(a)
|
||||
elif isinstance(a, FieldExpression):
|
||||
converted.append(str(a))
|
||||
else:
|
||||
converted.append(str(a))
|
||||
self._qs = self._qs.order_by(*converted)
|
||||
return self
|
||||
|
||||
def first(self):
|
||||
if self._cache_enabled:
|
||||
cache = _get_cache()
|
||||
key = self._cache_key('first')
|
||||
if key in cache:
|
||||
return cache[key]
|
||||
obj = self._qs.first()
|
||||
if obj is not None:
|
||||
_get_session().track(obj)
|
||||
if self._cache_enabled:
|
||||
_get_cache()[self._cache_key('first')] = obj
|
||||
return obj
|
||||
|
||||
def all(self):
|
||||
"""返回 self 以支持链式调用(与 Django QuerySet.all() 一致)。
|
||||
如需获取列表,使用 .to_list() 或 list()。"""
|
||||
self._qs = self._qs.all()
|
||||
return self
|
||||
|
||||
def to_list(self):
|
||||
"""执行查询并返回模型实例列表。"""
|
||||
if self._cache_enabled:
|
||||
cache = _get_cache()
|
||||
key = self._cache_key('to_list')
|
||||
if key in cache:
|
||||
return cache[key]
|
||||
results = list(self._qs)
|
||||
session = _get_session()
|
||||
for obj in results:
|
||||
session.track(obj)
|
||||
if self._cache_enabled:
|
||||
_get_cache()[self._cache_key('to_list')] = results
|
||||
return results
|
||||
|
||||
def count(self):
|
||||
if self._cache_enabled:
|
||||
cache = _get_cache()
|
||||
key = self._cache_key('count')
|
||||
if key in cache:
|
||||
return cache[key]
|
||||
result = self._qs.count()
|
||||
if self._cache_enabled:
|
||||
_get_cache()[self._cache_key('count')] = result
|
||||
return result
|
||||
|
||||
def get(self, *args, **kwargs):
|
||||
if args:
|
||||
obj = self._qs.filter(*args).get(**kwargs)
|
||||
else:
|
||||
obj = self._qs.get(**kwargs)
|
||||
_get_session().track(obj)
|
||||
return obj
|
||||
|
||||
def create(self, **kwargs):
|
||||
obj = self._model.objects.create(**kwargs)
|
||||
_invalidate_cache()
|
||||
return obj
|
||||
|
||||
def get_or_create(self, defaults=None, **kwargs):
|
||||
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):
|
||||
obj, created = self._model.objects.update_or_create(defaults=defaults, **kwargs)
|
||||
_invalidate_cache()
|
||||
_get_session().track(obj)
|
||||
return obj, created
|
||||
|
||||
def none(self):
|
||||
self._qs = self._qs.none()
|
||||
return self
|
||||
|
||||
def last(self):
|
||||
return self._qs.last()
|
||||
|
||||
def limit(self, n):
|
||||
self._qs = self._qs[:n]
|
||||
return self
|
||||
|
||||
def offset(self, n):
|
||||
self._qs = self._qs[n:]
|
||||
return self
|
||||
|
||||
def join(self, *args):
|
||||
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):
|
||||
self._qs = self._qs.select_related(*fields)
|
||||
return self
|
||||
|
||||
def prefetch_related(self, *fields):
|
||||
self._qs = self._qs.prefetch_related(*fields)
|
||||
return self
|
||||
|
||||
def only(self, *fields):
|
||||
self._qs = self._qs.only(*fields)
|
||||
return self
|
||||
|
||||
def defer(self, *fields):
|
||||
self._qs = self._qs.defer(*fields)
|
||||
return self
|
||||
|
||||
def select_for_update(self, nowait=False, skip_locked=False, of=(), no_key=False):
|
||||
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):
|
||||
result = self._model.objects.bulk_create(
|
||||
objs, batch_size=batch_size, ignore_conflicts=ignore_conflicts
|
||||
)
|
||||
_invalidate_cache()
|
||||
return result
|
||||
|
||||
def annotate(self, **kwargs):
|
||||
self._qs = self._qs.annotate(**kwargs)
|
||||
return self
|
||||
|
||||
def aggregate(self, **kwargs):
|
||||
return self._qs.aggregate(**kwargs)
|
||||
|
||||
def values(self, *fields):
|
||||
if not fields:
|
||||
# 无参数:返回 Django .values()(所有字段字典)
|
||||
return _ValuesResult(self._qs.values(), [], [])
|
||||
|
||||
annotations = {}
|
||||
converted = []
|
||||
for i, f in enumerate(fields):
|
||||
if isinstance(f, _AggregateExpr):
|
||||
annotations[f'val{i}'] = f.aggregate
|
||||
converted.append(f'val{i}')
|
||||
elif isinstance(f, (models.Aggregate, models.Expression)):
|
||||
annotations[f'val{i}'] = f
|
||||
converted.append(f'val{i}')
|
||||
elif isinstance(f, str):
|
||||
converted.append(f)
|
||||
else:
|
||||
converted.append(str(f))
|
||||
|
||||
qs = self._qs
|
||||
if annotations:
|
||||
qs = qs.annotate(**annotations)
|
||||
# 使用 Django 原生 .values() 返回字典
|
||||
qs = qs.values(*converted)
|
||||
return _ValuesResult(qs, converted, fields)
|
||||
|
||||
def with_entities(self, *fields):
|
||||
return self.values(*fields)
|
||||
|
||||
def params(self, **kwargs):
|
||||
return self
|
||||
|
||||
def values_list(self, *fields, flat=False):
|
||||
return self._qs.values_list(*fields, flat=flat)
|
||||
|
||||
def update(self, **kwargs):
|
||||
result = self._qs.update(**kwargs)
|
||||
_invalidate_cache()
|
||||
return result
|
||||
|
||||
def delete(self):
|
||||
result = self._qs.delete()
|
||||
_invalidate_cache()
|
||||
return result
|
||||
|
||||
def exists(self):
|
||||
return self._qs.exists()
|
||||
|
||||
def distinct(self):
|
||||
self._qs = self._qs.distinct()
|
||||
return self
|
||||
|
||||
def paginate(self, page=1, per_page=20, error_out=False):
|
||||
from django.core.paginator import Paginator
|
||||
if not self._qs.ordered:
|
||||
meta = getattr(self._model, '_meta', None)
|
||||
if meta and meta.pk:
|
||||
self._qs = self._qs.order_by(meta.pk.name)
|
||||
paginator = Paginator(self._qs, per_page)
|
||||
try:
|
||||
page_obj = paginator.page(page)
|
||||
except Exception:
|
||||
page_obj = paginator.page(1)
|
||||
page_obj.total = paginator.count
|
||||
page_obj.items = page_obj.object_list
|
||||
page_obj.pages = paginator.num_pages
|
||||
page_obj.page = page
|
||||
page_obj.per_page = per_page
|
||||
page_obj.has_next = page_obj.has_next()
|
||||
page_obj.has_prev = page_obj.has_previous()
|
||||
page_obj.next_num = page_obj.next_page_number() if page_obj.has_next else None
|
||||
page_obj.prev_num = page_obj.previous_page_number() if page_obj.has_prev else None
|
||||
return page_obj
|
||||
|
||||
def no_cache(self):
|
||||
self._cache_enabled = False
|
||||
return self
|
||||
|
||||
def using(self, alias):
|
||||
self._qs = self._qs.using(alias)
|
||||
return self
|
||||
|
||||
def __getitem__(self, key):
|
||||
"""支持切片 [start:end] 和索引 [n]"""
|
||||
if isinstance(key, slice):
|
||||
self._qs = self._qs[key]
|
||||
return self
|
||||
else:
|
||||
# 单个索引,返回模型实例
|
||||
return self._qs[key]
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self._qs)
|
||||
|
||||
def __len__(self):
|
||||
return self.count()
|
||||
|
||||
def __bool__(self):
|
||||
return self.exists()
|
||||
|
||||
def __repr__(self):
|
||||
return f'FluentQuery({self._model.__name__})'
|
||||
|
||||
|
||||
class Session:
|
||||
def __init__(self):
|
||||
self._pending_adds = []
|
||||
self._pending_deletes = []
|
||||
self._tracked = set()
|
||||
|
||||
def add(self, obj):
|
||||
self._pending_adds.append(obj)
|
||||
|
||||
def delete(self, obj):
|
||||
self._pending_deletes.append(obj)
|
||||
self._tracked.discard(obj)
|
||||
|
||||
def track(self, obj):
|
||||
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}
|
||||
with transaction.atomic():
|
||||
for obj in self._pending_deletes:
|
||||
obj.delete()
|
||||
for obj in self._pending_adds:
|
||||
obj.save()
|
||||
for obj in self._tracked:
|
||||
if obj.pk is not None and (obj.__class__, obj.pk) in delete_pks:
|
||||
continue
|
||||
try:
|
||||
obj.save()
|
||||
except Exception:
|
||||
pass
|
||||
self._pending_adds.clear()
|
||||
self._pending_deletes.clear()
|
||||
self._tracked.clear()
|
||||
_invalidate_cache()
|
||||
|
||||
def rollback(self):
|
||||
self._pending_adds.clear()
|
||||
self._pending_deletes.clear()
|
||||
self._tracked.clear()
|
||||
_invalidate_cache()
|
||||
|
||||
def query(self, *args):
|
||||
model = None
|
||||
aggregates = []
|
||||
for a in args:
|
||||
if isinstance(a, type) and issubclass(a, models.Model):
|
||||
model = a
|
||||
elif isinstance(a, _AggregateExpr):
|
||||
if a.model:
|
||||
model = a.model
|
||||
aggregates.append(a.aggregate)
|
||||
elif isinstance(a, (models.Aggregate, models.Expression)):
|
||||
aggregates.append(a)
|
||||
if not aggregates and model:
|
||||
return FluentQuery(model)
|
||||
if aggregates and model:
|
||||
fq = FluentQuery(model)
|
||||
return fq.values(*args)
|
||||
return None
|
||||
|
||||
def flush(self):
|
||||
self.commit()
|
||||
|
||||
def execute(self, stmt):
|
||||
if isinstance(stmt, _DeleteStmt):
|
||||
stmt.execute()
|
||||
elif isinstance(stmt, str):
|
||||
connection = default_db_connection
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(stmt)
|
||||
|
||||
|
||||
class DB:
|
||||
@property
|
||||
def session(self):
|
||||
return _get_session()
|
||||
|
||||
@staticmethod
|
||||
def or_(*args):
|
||||
result = None
|
||||
for q in args:
|
||||
if isinstance(q, Q):
|
||||
if result is None:
|
||||
result = q
|
||||
else:
|
||||
result |= q
|
||||
return result if result is not None else Q()
|
||||
|
||||
@staticmethod
|
||||
def and_(*args):
|
||||
result = Q()
|
||||
for q in args:
|
||||
if isinstance(q, Q):
|
||||
result &= q
|
||||
return result
|
||||
|
||||
|
||||
db = DB()
|
||||
|
||||
|
||||
def FQ(model_class):
|
||||
return FluentQuery(model_class)
|
||||
Reference in New Issue
Block a user