659 lines
20 KiB
Python
659 lines
20 KiB
Python
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 __getitem__(self, key):
|
||
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._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)
|