加入了 GVSDSDK 模块,进行了 QModel 兼容层的尝试,生产环境可用

This commit is contained in:
2026-06-16 00:13:10 +08:00
parent 9d9cfa53a4
commit 63e0f4edfc
68 changed files with 10256 additions and 1143 deletions

651
gvsdsdk/fluent.py Normal file
View 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)