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)