Files
Django/gvsdsdk/fluent.py

652 lines
20 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)