1
0
forked from erp-dev/erp
Files
erpnew/api_v2/views/ai.py

333 lines
12 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 logging
from django.conf import settings
from django.http import Http404
from rest_framework import authentication, mixins, serializers
from rest_framework.authentication import get_authorization_header
from rest_framework.exceptions import AuthenticationFailed
from rest_framework.generics import GenericAPIView
from rest_framework.permissions import IsAuthenticated
from rest_framework.response import Response
from rest_framework.views import APIView
from basic_info import models as basic_models
from flower.viewsets import LimitedLimitOffsetPagination
from printing.models import PrintingJob
from shipment.models import SalesItem, Shipment
logger = logging.getLogger(__name__)
class AgentAccessPrincipal:
id = None
username = "agent"
is_staff = False
is_superuser = False
is_anonymous = False
is_authenticated = True
class AgentAccessKeyAuthentication(authentication.BaseAuthentication):
"""
使用固定 AGENT_ACCESS_KEY 的轻量认证。
约定:
- 通过 Authorization 头直接传入密钥
- 不使用 Bearer 前缀
"""
def authenticate(self, request):
configured_key = str(getattr(settings, "AGENT_ACCESS_KEY", "") or "").strip()
if not configured_key:
logger.warning("Agent API authentication failed: AGENT_ACCESS_KEY is not configured")
raise AuthenticationFailed("AGENT_ACCESS_KEY 未配置")
provided_key = get_authorization_header(request).decode("utf-8").strip()
if not provided_key:
raise AuthenticationFailed("缺少 Authorization 请求头")
if provided_key != configured_key:
raise AuthenticationFailed("AGENT_ACCESS_KEY 无效")
return AgentAccessPrincipal(), None
def authenticate_header(self, request):
return "AgentAccessKey"
class AgentUnshippedShipmentQuerySerializer(serializers.Serializer):
merchant_id = serializers.IntegerField(min_value=1)
area = serializers.CharField(allow_blank=False, max_length=30)
limit = serializers.IntegerField(required=False, min_value=1)
offset = serializers.IntegerField(required=False, min_value=0)
def validate(self, attrs):
area = attrs["area"].strip()
if not area:
raise serializers.ValidationError({"area": "地区不能为空"})
merchant = basic_models.Merchant.objects.filter(id=attrs["merchant_id"]).first()
if merchant is None:
raise serializers.ValidationError({"merchant_id": "商户不存在"})
attrs["area"] = area
attrs["merchant"] = merchant
return attrs
class AgentSalesItemSerializer(serializers.ModelSerializer):
unit_display = serializers.CharField(source="get_unit_display", read_only=True)
printing_job_width = serializers.SerializerMethodField()
class Meta:
model = SalesItem
fields = [
"id",
"name",
"quantity",
"unit",
"unit_display",
"position",
"remark",
"printing_job_id",
"printing_job_width",
]
def get_printing_job_width(self, obj):
width_map = self.context.get("printing_job_width_map", {})
return width_map.get(obj.printing_job_id)
class AgentShipmentListSerializer(serializers.ModelSerializer):
customer_name = serializers.CharField(source="customer.name", read_only=True)
status_display = serializers.CharField(source="get_status_display", read_only=True)
fabric = serializers.SerializerMethodField()
order_description = serializers.SerializerMethodField()
delivery_id = serializers.IntegerField(read_only=True, allow_null=True)
sales_items = serializers.SerializerMethodField()
class Meta:
model = Shipment
fields = [
"id",
"merchant_id",
"customer",
"customer_name",
"fabric",
"order_description",
"shipment_date",
"address",
"contact_name",
"contact_phone",
"area",
"remark",
"status",
"status_display",
"external_id",
"geo_coordinates",
"delivery_id",
"delivery",
"sales_items",
"created_at",
"updated_at",
]
def get_fabric(self, obj):
rel = getattr(obj, "items", None)
if rel is None:
return None
first_item = next((item for item in obj.items.all() if item.delete_at is None), None)
if first_item is None or not first_item.printing_job_id:
return None
from printing.models import PrintingJob
printing_job = (
PrintingJob.objects.filter(id=first_item.printing_job_id)
.select_related("printing_order")
.first()
)
if printing_job is None or printing_job.printing_order is None:
return None
return printing_job.printing_order.fabric
def get_order_description(self, obj):
rel = getattr(obj, "items", None)
if rel is None:
return None
first_item = next((item for item in obj.items.all() if item.delete_at is None), None)
if first_item is None or not first_item.printing_job_id:
return None
from printing.models import PrintingJob
printing_job = (
PrintingJob.objects.filter(id=first_item.printing_job_id)
.select_related("printing_order")
.first()
)
if printing_job is None or printing_job.printing_order is None:
return None
return printing_job.printing_order.rolling_warn
def get_sales_items(self, obj):
items = [i for i in obj.items.all() if i.delete_at is None]
return AgentSalesItemSerializer(items, many=True, context=self.context).data
class AgentUnshippedShipmentListView(mixins.ListModelMixin, GenericAPIView):
"""
Agent 专用:查询尚未进入送货单的出货单列表。
GET /api/v2/ai/shipments/unshipped/?merchant_id=<id>&area=<地区>&limit=<n>&offset=<n>
"""
authentication_classes = [AgentAccessKeyAuthentication]
permission_classes = [IsAuthenticated]
serializer_class = AgentShipmentListSerializer
pagination_class = LimitedLimitOffsetPagination
def get(self, request):
query_serializer = AgentUnshippedShipmentQuerySerializer(data=request.query_params)
query_serializer.is_valid(raise_exception=True)
self.validated_query = query_serializer.validated_data
return self.list(request)
def list(self, request, *args, **kwargs):
queryset = self.filter_queryset(self.get_queryset())
page = self.paginate_queryset(queryset)
shipments = page if page is not None else queryset
# Batch-fetch printing job widths to avoid N+1
pj_ids = {
item.printing_job_id
for shipment in shipments
for item in shipment.items.all()
if item.printing_job_id and item.delete_at is None
}
width_map = {}
if pj_ids:
width_map = {
job.id: job.printing_order.width
for job in PrintingJob.objects.filter(id__in=pj_ids).select_related("printing_order")
}
ctx = {**self.get_serializer_context(), "printing_job_width_map": width_map}
serializer = self.get_serializer(shipments, many=True, context=ctx)
if page is not None:
return self.get_paginated_response(serializer.data)
return Response(serializer.data)
def get_queryset(self):
validated_query = getattr(self, "validated_query", None)
if validated_query is None:
query_serializer = AgentUnshippedShipmentQuerySerializer(data=self.request.query_params)
query_serializer.is_valid(raise_exception=True)
validated_query = query_serializer.validated_data
merchant = validated_query["merchant"]
area = validated_query["area"]
return (
Shipment.objects.filter(
merchant=merchant,
area=area,
delivery__isnull=True,
)
.select_related("customer", "merchant")
.prefetch_related("items")
.order_by("-created_at", "-id")
)
# ---------------------------------------------------------------------------
# Vehicle views
# ---------------------------------------------------------------------------
class AgentMerchantQuerySerializer(serializers.Serializer):
"""Base serializer that validates merchant_id and returns merchant instance."""
merchant_id = serializers.IntegerField(min_value=1)
def validate_merchant_id(self, value):
merchant = basic_models.Merchant.objects.filter(id=value).first()
if merchant is None:
raise serializers.ValidationError("商户不存在")
return merchant
class AgentTransportVehicleCapacitySerializer(serializers.ModelSerializer):
class Meta:
model = basic_models.TransportVehicleMaterialCapacity
fields = ["id", "material_name", "capacity"]
class AgentTransportVehicleSerializer(serializers.ModelSerializer):
material_capacities = AgentTransportVehicleCapacitySerializer(many=True, read_only=True)
class Meta:
model = basic_models.TransportVehicle
fields = ["id", "merchant_id", "name", "license_plate", "material_capacities", "created_at", "updated_at"]
class AgentTransportVehicleListView(mixins.ListModelMixin, GenericAPIView):
"""
Agent 专用:查询商户下的所有运输车辆及其物料容量。
GET /api/v2/ai/transport-vehicles/?merchant_id=<id>
"""
authentication_classes = [AgentAccessKeyAuthentication]
permission_classes = [IsAuthenticated]
serializer_class = AgentTransportVehicleSerializer
pagination_class = LimitedLimitOffsetPagination
def get(self, request):
query_serializer = AgentMerchantQuerySerializer(data=request.query_params)
query_serializer.is_valid(raise_exception=True)
self.merchant = query_serializer.validated_data["merchant_id"]
return self.list(request)
def get_queryset(self):
merchant = getattr(self, "merchant", None)
if merchant is None:
query_serializer = AgentMerchantQuerySerializer(data=self.request.query_params)
query_serializer.is_valid(raise_exception=True)
merchant = query_serializer.validated_data["merchant_id"]
return (
basic_models.TransportVehicle.objects.filter(merchant=merchant)
.prefetch_related("material_capacities")
.order_by("id")
)
class AgentTransportVehicleDetailView(APIView):
"""
Agent 专用:通过车牌号查询运输车辆详情。
GET /api/v2/ai/transport-vehicles/<license_plate>/?merchant_id=<id>
license_plate 在同一商户内唯一unique_together: merchant + license_plate
"""
authentication_classes = [AgentAccessKeyAuthentication]
permission_classes = [IsAuthenticated]
def get(self, request, license_plate):
query_serializer = AgentMerchantQuerySerializer(data=request.query_params)
query_serializer.is_valid(raise_exception=True)
merchant = query_serializer.validated_data["merchant_id"]
try:
vehicle = (
basic_models.TransportVehicle.objects
.prefetch_related("material_capacities")
.get(merchant=merchant, license_plate=license_plate)
)
except basic_models.TransportVehicle.DoesNotExist:
raise Http404
return Response(AgentTransportVehicleSerializer(vehicle).data)