如何在Django REST Framework中使用带参数的原生SQL查询?附示例代码
在Django REST Framework中实现带参数的原生SQL聚合查询
我来帮你一步步搞定这个需求,咱们直接从代码修改入手:
1. 创建自定义序列化器(适配查询结果)
你的原生SQL返回的是聚合后的统计数据(date、driver_id、TRIP_ADJUSTMENT),和原模型的字段不完全匹配,所以需要新建一个序列化器来处理这个结果:
# serializer.py from rest_framework import serializers class DriverPenaltyStatsSerializer(serializers.Serializer): date = serializers.DateField() driver_id = serializers.IntegerField() TRIP_ADJUSTMENT = serializers.DecimalField(max_digits=10, decimal_places=2)
2. 修改视图集,添加自定义查询接口
在DriverPenaltyViewSet里新增一个自定义action,用来处理这个统计查询。我们用Django的connection执行原生SQL,能更灵活地处理聚合逻辑:
# views.py from django.db import connection from rest_framework import viewsets, status from rest_framework.decorators import action from rest_framework.response import Response from .models import DriverPenalty from .serializers import DriverPenaltySerializer, DriverPenaltyStatsSerializer from rest_framework.permissions import IsAuthenticated class DriverPenaltyViewSet(viewsets.ModelViewSet): permission_classes=(IsAuthenticated,) queryset= DriverPenalty.objects.all().order_by('driver_id') serializer_class=DriverPenaltySerializer @action(detail=False, methods=['get']) def trip_adjustment_stats(self, request): # 从请求参数中获取查询条件 driver_id = request.query_params.get('driver_id') start_date = request.query_params.get('start_date') end_date = request.query_params.get('end_date') # 校验必要参数是否齐全 if not all([driver_id, start_date, end_date]): return Response( {"error": "缺少必要参数:driver_id、start_date、end_date"}, status=status.HTTP_400_BAD_REQUEST ) # 编写原生SQL语句(注意加上GROUP BY,否则会合并所有结果为一行) sql = """ SELECT date, driver_id, sum(case when type='TRIP_ADJUSTMENT' then amount ELSE 0 end) as TRIP_ADJUSTMENT from fleet_driverpenalty WHERE driver_id= %s and (date BETWEEN %s and %s) GROUP BY date, driver_id """ # 执行SQL查询并处理结果 with connection.cursor() as cursor: cursor.execute(sql, [driver_id, start_date, end_date]) # 获取字段名,把结果转成字典格式 columns = [col[0] for col in cursor.description] results = [dict(zip(columns, row)) for row in cursor.fetchall()] # 序列化结果并返回 serializer = DriverPenaltyStatsSerializer(results, many=True) return Response(serializer.data)
这里有几个关键细节要注意:
- 一定要加上
GROUP BY date, driver_id,否则sum会把所有符合条件的记录合并成一条,无法按日期和司机分组统计 - 用列表传递参数给
cursor.execute(),Django会自动处理SQL注入防护,绝对不要手动拼接SQL字符串 @action(detail=False)表示这是一个列表级接口(不需要传入模型主键),请求方法设为get符合统计查询的场景
3. 测试接口
启动服务后,你可以通过类似下面的URL访问统计接口:
GET /api/driverpenalty/trip_adjustment_stats/?driver_id=1&start_date=2024-01-01&end_date=2024-06-30
替换/api/driverpenalty/为你的实际路由前缀,参数换成真实的司机ID和日期范围即可。
可选:用RawQuerySet的方式(如果需要关联模型)
如果你的查询结果需要和原模型实例关联,也可以尝试DriverPenalty.objects.raw(),不过因为这里是聚合查询,返回的不是完整模型实例,所以前面用connection的方式更合适。如果要试的话,代码大概是这样:
raw_query = DriverPenalty.objects.raw(sql, [driver_id, start_date, end_date]) # 注意:raw()返回的模型实例中,TRIP_ADJUSTMENT会作为额外属性存在
内容的提问来源于stack exchange,提问作者Ashish Pandey
相关产品推荐
相关产品推荐

