From ddb3cc6f3f7296f115cd6b15226c9bbda5f864c9 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Tue, 4 Aug 2026 13:42:59 +0800 Subject: [PATCH 01/15] fix(wpm): derive number date filters from rule --- apps/wpm/tests/test_number_rule.py | 64 ++++++++++++++++++++++++++++++ apps/wpm/views.py | 18 ++++++--- 2 files changed, 77 insertions(+), 5 deletions(-) create mode 100644 apps/wpm/tests/test_number_rule.py diff --git a/apps/wpm/tests/test_number_rule.py b/apps/wpm/tests/test_number_rule.py new file mode 100644 index 00000000..ee8e7d5e --- /dev/null +++ b/apps/wpm/tests/test_number_rule.py @@ -0,0 +1,64 @@ +from datetime import date +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from django.test import SimpleTestCase + +from apps.wpm.views import MlogbInViewSet + + +class GenNumberWithRuleFilterTests(SimpleTestCase): + def test_date_filters_follow_rule_placeholders(self): + cases = [ + ("SN-{n_count:04d}", {}), + ("{c_year}-{n_count:04d}", {"year": 2026}), + ("{c_year2}-{n_count:04d}", {"year": 2026}), + ("{c_month:02d}-{n_count:04d}", {"month": 8}), + ("{c_day:02d}-{n_count:04d}", {"day": 4}), + ( + "{c_year}{c_month:02d}{c_day:02d}-{n_count:04d}", + {"year": 2026, "month": 8, "day": 4}, + ), + ] + material = SimpleNamespace(model=None) + mlog = SimpleNamespace( + handle_date=date(2026, 8, 4), + mgroup=SimpleNamespace(process=SimpleNamespace(id=123)), + ) + + for rule, expected_dates in cases: + with self.subTest(rule=rule): + queryset = MagicMock() + queryset.annotate.return_value = queryset + queryset.order_by.return_value = queryset + queryset.last.return_value = None + + with patch("apps.wpmw.models.Wpr.objects.filter", return_value=queryset) as mock_filter: + MlogbInViewSet.gen_number_with_rule(rule, material, mlog) + + filters = mock_filter.call_args.kwargs + date_prefix = "wpr_mlogbw__mlogb__mlog__handle_date__" + actual_dates = { + key.removeprefix(date_prefix): value + for key, value in filters.items() + if key.startswith(date_prefix) + } + self.assertEqual(actual_dates, expected_dates) + + def test_escaped_date_placeholder_text_does_not_add_filter(self): + queryset = MagicMock() + queryset.annotate.return_value = queryset + queryset.order_by.return_value = queryset + queryset.last.return_value = None + material = SimpleNamespace(model=None) + mlog = SimpleNamespace( + handle_date=date(2026, 8, 4), + mgroup=SimpleNamespace(process=SimpleNamespace(id=123)), + ) + + with patch("apps.wpmw.models.Wpr.objects.filter", return_value=queryset) as mock_filter: + MlogbInViewSet.gen_number_with_rule("{{c_year}}-{n_count:04d}", material, mlog) + + self.assertFalse( + any("handle_date" in key for key in mock_filter.call_args.kwargs) + ) diff --git a/apps/wpm/views.py b/apps/wpm/views.py index d7f01582..60e747ed 100644 --- a/apps/wpm/views.py +++ b/apps/wpm/views.py @@ -1,5 +1,6 @@ import math import re +from string import Formatter from django.db import transaction from rest_framework.decorators import action @@ -1010,13 +1011,18 @@ class MlogbInViewSet(BulkCreateModelMixin, BulkUpdateModelMixin, BulkDestroyMode def gen_number_with_rule(cls, rule, material_out: Material, mlog: Mlog, gen_count=1): from apps.wpmw.models import Wpr + rule_fields = { + field_name + for _, field_name, _, _ in Formatter().parse(rule) + if field_name + } handle_date = mlog.handle_date c_year = handle_date.year c_year2 = str(c_year)[-2:] c_month = handle_date.month c_day = handle_date.day m_model = material_out.model - if 'm_model' in rule: + if "m_model" in rule_fields: if m_model is None: raise ParseError("生成编号出错:产品型号不能为空") elif m_model and m_model.islower(): @@ -1029,16 +1035,18 @@ class MlogbInViewSet(BulkCreateModelMixin, BulkUpdateModelMixin, BulkDestroyMode if connection.vendor == "postgresql" and connection.in_atomic_block: with connection.cursor() as cursor: cursor.execute("SELECT pg_advisory_xact_lock(hashtextextended(%s, 0))", [f"wpr_number_rule:{process.id}"]) - # 按生产日志查询, 流水号归零周期跟随规则中最细的日期占位符 + # 只按规则中实际使用的日期占位符筛选历史编号 wpr_filter = { "wpr_mlogbw__mlogb__material_out__isnull": False, "wpr_mlogbw__mlogb__mlog__mgroup__process": process, "wpr_mlogbw__mlogb__mlog__is_fix": False, "wpr_mlogbw__mlogb__mlog__submit_time__isnull": False, - "wpr_mlogbw__mlogb__mlog__handle_date__year": c_year, - "wpr_mlogbw__mlogb__mlog__handle_date__month": c_month, } - if "c_day" in rule: + if rule_fields & {"c_year", "c_year2"}: + wpr_filter["wpr_mlogbw__mlogb__mlog__handle_date__year"] = c_year + if "c_month" in rule_fields: + wpr_filter["wpr_mlogbw__mlogb__mlog__handle_date__month"] = c_month + if "c_day" in rule_fields: wpr_filter["wpr_mlogbw__mlogb__mlog__handle_date__day"] = c_day wpr = ( Wpr.objects.filter(**wpr_filter) From 69b5346031aedda54ff9e29b2ec75675b8b8f63d Mon Sep 17 00:00:00 2001 From: caoqianming Date: Tue, 4 Aug 2026 14:32:24 +0800 Subject: [PATCH 02/15] feat(wpm): sync wpr number from output edits --- apps/wpm/serializers.py | 8 +++ apps/wpm/tests/test_mlogbw_number.py | 80 ++++++++++++++++++++++++++++ apps/wpmw/models.py | 13 +++++ apps/wpmw/views.py | 9 +--- 4 files changed, 102 insertions(+), 8 deletions(-) create mode 100644 apps/wpm/tests/test_mlogbw_number.py diff --git a/apps/wpm/serializers.py b/apps/wpm/serializers.py index fc7fa8e9..53d4daad 100644 --- a/apps/wpm/serializers.py +++ b/apps/wpm/serializers.py @@ -983,10 +983,18 @@ class MlogbwCreateUpdateSerializer(CustomModelSerializer): mlogbw = self.save_ftest(mlogbw, ftest_data) return mlogbw + @transaction.atomic def update(self, instance, validated_data): + old_number = instance.number validated_data.pop("mlogb") ftest_data = validated_data.pop("ftest", None) mlogbw:Mlogbw = super().update(instance, validated_data) + if ( + mlogbw.number != old_number + and mlogbw.mlogb.material_out_id is not None + and mlogbw.wpr is not None + ): + mlogbw.wpr.change_number(mlogbw.number) if ftest_data: mlogbw = self.save_ftest(mlogbw, ftest_data) elif ftest_data is None: diff --git a/apps/wpm/tests/test_mlogbw_number.py b/apps/wpm/tests/test_mlogbw_number.py new file mode 100644 index 00000000..d6720918 --- /dev/null +++ b/apps/wpm/tests/test_mlogbw_number.py @@ -0,0 +1,80 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from django.test import SimpleTestCase + +from apps.wpm.serializers import MlogbwCreateUpdateSerializer +from apps.wpmw.models import Wpr + + +class MlogbwNumberUpdateTests(SimpleTestCase): + @staticmethod + def _apply_update(instance, validated_data): + instance.number = validated_data["number"] + return instance + + @patch("apps.wpm.serializers.CustomModelSerializer.update") + def test_output_number_update_changes_linked_wpr_number(self, base_update): + base_update.side_effect = self._apply_update + wpr = MagicMock() + instance = SimpleNamespace( + number="OLD-001", + mlogb=SimpleNamespace(material_out_id="material-out"), + wpr=wpr, + ftest=None, + ) + + MlogbwCreateUpdateSerializer.update.__wrapped__( + MlogbwCreateUpdateSerializer(), + instance, + {"mlogb": instance.mlogb, "number": "NEW-001"}, + ) + + wpr.change_number.assert_called_once_with("NEW-001") + + @patch("apps.wpm.serializers.CustomModelSerializer.update") + def test_input_number_update_does_not_change_wpr_number(self, base_update): + base_update.side_effect = self._apply_update + wpr = MagicMock() + instance = SimpleNamespace( + number="OLD-001", + mlogb=SimpleNamespace(material_out_id=None), + wpr=wpr, + ftest=None, + ) + + MlogbwCreateUpdateSerializer.update.__wrapped__( + MlogbwCreateUpdateSerializer(), + instance, + {"mlogb": instance.mlogb, "number": "NEW-001"}, + ) + + wpr.change_number.assert_not_called() + + @patch("apps.wpmw.models.MIOItemw.objects.filter") + @patch("apps.wpmw.models.Handoverbw.objects.filter") + @patch("apps.wpmw.models.Mlogbw.objects.filter") + @patch("apps.wpmw.models.Wpr.objects.filter") + def test_wpr_number_change_updates_all_number_copies( + self, + wpr_filter, + mlogbw_filter, + handoverbw_filter, + mioitemw_filter, + ): + conflict_qs = MagicMock() + conflict_qs.exists.return_value = False + current_qs = MagicMock() + wpr_filter.side_effect = [conflict_qs, current_qs] + mlogbw_qs = mlogbw_filter.return_value + handoverbw_qs = handoverbw_filter.return_value + mioitemw_qs = mioitemw_filter.return_value + wpr = Wpr(id="wpr-id", number="OLD-001") + + wpr.change_number("NEW-001") + + current_qs.update.assert_called_once_with(number="NEW-001") + mlogbw_qs.update.assert_called_once_with(number="NEW-001") + handoverbw_qs.update.assert_called_once_with(number="NEW-001") + mioitemw_qs.update.assert_called_once_with(number="NEW-001") + self.assertEqual(wpr.number, "NEW-001") diff --git a/apps/wpmw/models.py b/apps/wpmw/models.py index a655e2e0..bc5832ba 100644 --- a/apps/wpmw/models.py +++ b/apps/wpmw/models.py @@ -33,6 +33,19 @@ class Wpr(BaseModel): data = models.JSONField(verbose_name="数据", default=dict, blank=True) pre_info = models.JSONField(verbose_name="预处理信息", default=dict, blank=True, null=True) + def change_number(self, new_number): + """修改产品编号,并同步所有保存了编号副本的关联明细。""" + if self.number == new_number: + return + if Wpr.objects.filter(number=new_number).exists(): + raise ParseError("新编号已存在,不可使用") + + Wpr.objects.filter(id=self.id).update(number=new_number) + Mlogbw.objects.filter(wpr=self).update(number=new_number) + Handoverbw.objects.filter(wpr=self).update(number=new_number) + MIOItemw.objects.filter(wpr=self).update(number=new_number) + self.number = new_number + @classmethod def change_or_new( cls, wpr=None, number=None, mb=None, wm=None, old_mb=None, diff --git a/apps/wpmw/views.py b/apps/wpmw/views.py index 60c51b69..e72a989c 100644 --- a/apps/wpmw/views.py +++ b/apps/wpmw/views.py @@ -63,15 +63,8 @@ class WprViewSet(BulkUpdateModelMixin, CustomListModelMixin, CustomRetrieveModel vdata = sr.validated_data new_number = vdata["new_number"] old_number = vdata["old_number"] - if Wpr.objects.filter(number=new_number).exists(): - raise ParseError("新编号已存在,不可使用") wpr = Wpr.objects.get(number=old_number) - from apps.wpm.models import Mlogbw, Handoverbw - from apps.inm.models import MIOItemw - Wpr.objects.filter(id=wpr.id).update(number=new_number) - Mlogbw.objects.filter(wpr=wpr).update(number=new_number) - Handoverbw.objects.filter(wpr=wpr).update(number=new_number) - MIOItemw.objects.filter(wpr=wpr).update(number=new_number) + wpr.change_number(new_number) return Response() @action(methods=["post"], detail=False, perms_map={"post": "*"}, serializer_class=WprNewSerializer) From 5fb179eb9e678877be785388ef8ebc8c05640b19 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Tue, 4 Aug 2026 16:33:27 +0800 Subject: [PATCH 03/15] fix(develop): restrict unsafe debug endpoints --- apps/develop/tests.py | 33 +++++++++++++++++++++++++++++++-- apps/develop/urls.py | 11 ++++++++--- apps/develop/views.py | 27 +++++++++++++-------------- 3 files changed, 52 insertions(+), 19 deletions(-) diff --git a/apps/develop/tests.py b/apps/develop/tests.py index 7ce503c2..9a7f0fe3 100755 --- a/apps/develop/tests.py +++ b/apps/develop/tests.py @@ -1,3 +1,32 @@ -from django.test import TestCase +from django.test import SimpleTestCase +from rest_framework.permissions import IsAdminUser +from rest_framework.test import APIRequestFactory -# Create your tests here. +from apps.develop.views import ServerTime, TestViewSet + + +class DevelopApiPermissionTests(SimpleTestCase): + def setUp(self): + self.factory = APIRequestFactory() + + def test_test_endpoint_rejects_anonymous_requests(self): + request = self.factory.post( + '/api/develop/test/send_sms/', + {}, + format='json', + ) + + response = TestViewSet.as_view({'post': 'send_sms'})(request) + + self.assertIn(response.status_code, (401, 403)) + + def test_server_time_requires_admin(self): + request = self.factory.get('/api/develop/server_time/') + + response = ServerTime.as_view()(request) + + self.assertIn(response.status_code, (401, 403)) + + def test_develop_views_use_admin_permission(self): + self.assertEqual(TestViewSet.permission_classes, [IsAdminUser]) + self.assertEqual(ServerTime.permission_classes, [IsAdminUser]) diff --git a/apps/develop/urls.py b/apps/develop/urls.py index bd894ba8..1986082d 100755 --- a/apps/develop/urls.py +++ b/apps/develop/urls.py @@ -1,4 +1,5 @@ -from django.urls import path, include +from django.conf import settings +from django.urls import include, path from apps.develop.views import (BackupDatabase, BackupMedia, ReloadClientGit, ReloadServerGit, ReloadServerOnly, TestViewSet, CorrectViewSet, testScanHtml, ServerTime) from rest_framework.routers import DefaultRouter @@ -6,9 +7,11 @@ from rest_framework.routers import DefaultRouter API_BASE_URL = 'api/develop/' HTML_BASE_URL = 'dhtml/develop/' router = DefaultRouter() -router.register('test', TestViewSet, basename='api_test') router.register('correct', CorrectViewSet, basename='correct') +if settings.DEBUG: + router.register('test', TestViewSet, basename='api_test') + urlpatterns = [ path(API_BASE_URL + 'reload_server_git/', ReloadServerGit.as_view()), # path(API_BASE_URL + 'reload_web_git/', ReloadClientGit.as_view()), @@ -17,5 +20,7 @@ urlpatterns = [ path(API_BASE_URL + 'backup_media/', BackupMedia.as_view()), path(API_BASE_URL + 'server_time/', ServerTime.as_view()), path(API_BASE_URL, include(router.urls)), - path(HTML_BASE_URL + "testscan/", testScanHtml) ] + +if settings.DEBUG: + urlpatterns.append(path(HTML_BASE_URL + "testscan/", testScanHtml)) diff --git a/apps/develop/views.py b/apps/develop/views.py index 03e493da..0c4c2fbe 100755 --- a/apps/develop/views.py +++ b/apps/develop/views.py @@ -2,7 +2,7 @@ from rest_framework.views import APIView from rest_framework.exceptions import ParseError -from rest_framework.permissions import IsAdminUser, AllowAny +from rest_framework.permissions import IsAdminUser from rest_framework.response import Response from rest_framework.serializers import Serializer from rest_framework.decorators import action @@ -40,11 +40,7 @@ from datetime import datetime # Create your views here. class ServerTime(APIView): - - def get_permissions(self): - if self.request.method == 'GET': - return [AllowAny()] - return [IsAdminUser()] + permission_classes = [IsAdminUser] @swagger_auto_schema(responses={200: ServerTimeSerializer}) def get(self, request): @@ -62,9 +58,13 @@ class ServerTime(APIView): 修改服务器时间 """ - command = f'date -s "{request.data["server_time"]}"' + serializer = ServerTimeSerializer(data=request.data) + serializer.is_valid(raise_exception=True) + server_time = serializer.validated_data['server_time'].strftime( + "%Y-%m-%d %H:%M:%S" + ) completed = subprocess.run( - ["sudo", "-S", "sh", "-c", command], # 添加 -S 参数 + ["sudo", "-S", "date", "-s", server_time], input=SD_PWD + "\n", # 注意要在密码后加换行符 capture_output=True, text=True @@ -269,10 +269,9 @@ class CorrectViewSet(CustomGenericViewSet): class TestViewSet(CustomGenericViewSet): perms_map = {} - authentication_classes = () - permission_classes = () + permission_classes = [IsAdminUser] - @action(methods=['post'], detail=False, serializer_class=SendSmsSerializer, authentication_classes=()) + @action(methods=['post'], detail=False, serializer_class=SendSmsSerializer) def send_sms(self, request, pk=None): """发送短信测试 @@ -565,7 +564,7 @@ class TestViewSet(CustomGenericViewSet): # correct_card_time() # return Response() - @action(methods=['post'], detail=False, serializer_class=Serializer, permission_classes=[]) + @action(methods=['post'], detail=False, serializer_class=Serializer) @transaction.atomic def correct_data(self, request, pk=None): """修正数据 @@ -674,7 +673,7 @@ class TestViewSet(CustomGenericViewSet): Ticket.objects.get_queryset(all=True).delete() return Response() - @action(methods=['post'], detail=False, serializer_class=Serializer, permission_classes=[]) + @action(methods=['post'], detail=False, serializer_class=Serializer) def test_cal(self, request, pk=None): from apps.wpm.tasks import cal_exp_duration_sec cal_exp_duration_sec('3397169058570170368') @@ -710,4 +709,4 @@ html_str = """ """ def testScanHtml(request): - return HttpResponse(html_str) \ No newline at end of file + return HttpResponse(html_str) From 4bf5f1e58513e51eb6aac6505db6748aa643264b Mon Sep 17 00:00:00 2001 From: caoqianming Date: Wed, 5 Aug 2026 15:01:44 +0800 Subject: [PATCH 04/15] feat(bi): expose dataset query workflow --- apps/bi/serializers.py | 39 ++++++++++++++++++++++-- apps/bi/tests.py | 26 ++++++++++++++-- apps/bi/views.py | 69 ++++++++++++++++++++++++++++++++++++++++-- apps/utils/mixins.py | 28 +++++++++++++---- apps/wpm/views.py | 1 + 5 files changed, 150 insertions(+), 13 deletions(-) diff --git a/apps/bi/serializers.py b/apps/bi/serializers.py index 35d8fe1b..5a666dc5 100644 --- a/apps/bi/serializers.py +++ b/apps/bi/serializers.py @@ -18,10 +18,35 @@ class DatasetCreateUpdateSerializer(CustomModelSerializer): class DatasetSerializer(CustomModelSerializer): + description = serializers.CharField( + label="适用场景与统计口径", + help_text="说明该数据集适合回答的问题、指标口径、参数格式和返回字段含义", + required=False, + allow_blank=True, + ) + default_param = serializers.JSONField( + label="默认查询参数", + help_text="执行时可覆盖的参数及默认值;内部 SQL 片段参数应保留默认值", + required=False, + ) + test_param = serializers.JSONField( + label="测试查询参数", + help_text="数据集维护时使用的示例参数,普通查询优先参考 description", + required=False, + ) + class Meta: model = Dataset fields = '__all__' + +class DatasetListResponseSerializer(serializers.Serializer): + count = serializers.IntegerField(label="数据集总数") + next = serializers.URLField(required=False, allow_null=True) + previous = serializers.URLField(required=False, allow_null=True) + results = DatasetSerializer(many=True) + + class DatasetRecordSerializer(CustomModelSerializer): class Meta: model = DatasetRecord @@ -36,6 +61,14 @@ class DatasetRecordSerializer(CustomModelSerializer): class DataExecSerializer(serializers.Serializer): query = serializers.JSONField( - label="查询字典参数", required=False, allow_null=True) - is_test = serializers.BooleanField(label='是否测试', default=False) - raise_exception = serializers.BooleanField(label='是否直接报错', default=False) + label="查询字典参数", + help_text="按所选数据集 description/default_param 声明的业务参数填写", + required=False, + allow_null=True, + ) + is_test = serializers.BooleanField( + label='是否测试', help_text="普通业务查询固定为 false", default=False + ) + raise_exception = serializers.BooleanField( + label='是否直接报错', help_text="建议为 true,便于修正缺失或非法参数", default=True + ) diff --git a/apps/bi/tests.py b/apps/bi/tests.py index 7ce503c2..30d14569 100644 --- a/apps/bi/tests.py +++ b/apps/bi/tests.py @@ -1,3 +1,25 @@ -from django.test import TestCase +from django.test import SimpleTestCase -# Create your tests here. +from apps.bi.serializers import ( + DataExecSerializer, + DatasetListResponseSerializer, + DatasetSerializer, +) +from apps.bi.views import DatasetViewSet + + +class DatasetAgentDiscoveryTests(SimpleTestCase): + def test_dataset_catalog_searches_description(self): + self.assertEqual( + DatasetViewSet.search_fields, + ["name", "code", "description"], + ) + + def test_dataset_schema_explains_discovery_and_exec_parameters(self): + dataset = DatasetSerializer() + execute = DataExecSerializer() + + self.assertIn("适合回答的问题", dataset.fields["description"].help_text) + self.assertIn("业务参数", execute.fields["query"].help_text) + self.assertTrue(execute.fields["raise_exception"].default) + self.assertIn("results", DatasetListResponseSerializer().fields) diff --git a/apps/bi/views.py b/apps/bi/views.py index 6912ae74..7c5ac1d3 100644 --- a/apps/bi/views.py +++ b/apps/bi/views.py @@ -3,7 +3,13 @@ from apps.utils.viewsets import CustomModelViewSet, CustomGenericViewSet from rest_framework.decorators import action from rest_framework.response import Response from apps.bi.models import Dataset, DatasetRecord -from apps.bi.serializers import DatasetSerializer, DatasetCreateUpdateSerializer, DataExecSerializer, DatasetRecordSerializer +from apps.bi.serializers import ( + DataExecSerializer, + DatasetCreateUpdateSerializer, + DatasetListResponseSerializer, + DatasetRecordSerializer, + DatasetSerializer, +) from django.apps import apps import concurrent.futures from django.core.cache import cache @@ -13,6 +19,8 @@ from rest_framework.exceptions import ParseError from rest_framework.generics import get_object_or_404 from apps.utils.mixins import ListModelMixin import logging +from drf_yasg import openapi +from drf_yasg.utils import swagger_auto_schema myLogger = logging.getLogger('log') # Create your views here. @@ -22,9 +30,54 @@ class DatasetViewSet(CustomModelViewSet): serializer_class = DatasetSerializer create_serializer_class = DatasetCreateUpdateSerializer update_serializer_class = DatasetCreateUpdateSerializer - search_fields = ['name', 'code'] + search_fields = ['name', 'code', 'description'] ordering = ['name', 'code', 'id'] + @swagger_auto_schema( + operation_id="bi_dataset_list", + operation_summary="查询复杂统计报表的数据集目录", + operation_description=( + "产量、良率、缺陷、库存、绩效、趋势和按日/月汇总等统计聚合查询的统一入口。" + "先调用本接口,根据 name、description、default_param 和 test_param 选择数据集," + "再调用 bi_dataset_exec。建议使用 query={id,name,code,description,default_param," + "test_param,enabled} 裁剪字段,并设置 page_size=100 查看完整目录;" + "search 可按名称、code 或 description 检索。" + ), + manual_parameters=[ + openapi.Parameter( + "search", + openapi.IN_QUERY, + description="按数据集名称、code 或适用场景关键词检索", + type=openapi.TYPE_STRING, + ), + openapi.Parameter( + "page", + openapi.IN_QUERY, + description="页码,从 1 开始", + type=openapi.TYPE_INTEGER, + ), + openapi.Parameter( + "page_size", + openapi.IN_QUERY, + description="每页数量;当前目录建议传 100", + type=openapi.TYPE_INTEGER, + ), + openapi.Parameter( + "query", + openapi.IN_QUERY, + description=( + "django-restql 字段裁剪表达式,例如 " + "{id,name,code,description,default_param,test_param,enabled}" + ), + type=openapi.TYPE_STRING, + ), + ], + responses={200: DatasetListResponseSerializer}, + tags=["BI 数据集与报表"], + ) + def list(self, request, *args, **kwargs): + return super().list(request, *args, **kwargs) + def get_object(self): """ Returns the object the view is displaying. @@ -57,6 +110,18 @@ class DatasetViewSet(CustomModelViewSet): return obj + @swagger_auto_schema( + operation_id="bi_dataset_exec", + operation_summary="执行已配置的只读统计数据集", + operation_description=( + "使用 dataset list 返回的 id 或 code 执行数据集。body.query 只填写该数据集" + "description/default_param 声明的业务参数;正常查询设置 is_test=false。" + "统计聚合使用本接口,日志和业务明细列表用于逐条追溯。" + ), + request_body=DataExecSerializer, + responses={200: DatasetSerializer}, + tags=["BI 数据集与报表"], + ) @action(methods=['post'], detail=True, perms_map={'post': 'dataset.exec'}, serializer_class=DataExecSerializer, cache_seconds=0, logging_methods=[]) def exec(self, request, pk=None): """执行sql查询 diff --git a/apps/utils/mixins.py b/apps/utils/mixins.py index 990d5926..61382ded 100755 --- a/apps/utils/mixins.py +++ b/apps/utils/mixins.py @@ -207,12 +207,28 @@ class CustomRetrieveModelMixin(RetrieveModelMixin): class CustomListModelMixin(ListModelMixin): - @swagger_auto_schema(manual_parameters=[ - openapi.Parameter(name="query", in_=openapi.IN_QUERY, description="定制返回数据", - type=openapi.TYPE_STRING, required=False), - openapi.Parameter(name="with_children", in_=openapi.IN_QUERY, description="带有children(yes/no/count)", - type=openapi.TYPE_STRING, required=False), - ]) + @swagger_auto_schema( + operation_description=( + "通用列表接口用于记录或目录浏览以及逐条追溯。跨时间范围的产量、良率、缺陷、" + "库存、绩效和趋势等统计聚合,优先查询 BI dataset 目录并执行匹配的数据集。" + ), + manual_parameters=[ + openapi.Parameter( + name="query", + in_=openapi.IN_QUERY, + description="django-restql 返回字段裁剪表达式", + type=openapi.TYPE_STRING, + required=False, + ), + openapi.Parameter( + name="with_children", + in_=openapi.IN_QUERY, + description="带有children(yes/no/count)", + type=openapi.TYPE_STRING, + required=False, + ), + ], + ) def list(self, request, *args, **kwargs): queryset = self.filter_queryset(self.get_queryset()) diff --git a/apps/wpm/views.py b/apps/wpm/views.py index 60e747ed..1084145f 100644 --- a/apps/wpm/views.py +++ b/apps/wpm/views.py @@ -354,6 +354,7 @@ class MlogViewSet(CustomModelViewSet): return super().get_serializer_class() @swagger_auto_schema( + operation_summary="查询生产日志明细(逐条追溯)", manual_parameters=[ openapi.Parameter(name="query", in_=openapi.IN_QUERY, description="定制返回数据", type=openapi.TYPE_STRING, required=False), openapi.Parameter(name="with_children", in_=openapi.IN_QUERY, description="带有children(yes/no/count)", type=openapi.TYPE_STRING, required=False), From ed952d2d3a1fee5c5b8390a8e1b7b210d8be1ad7 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Wed, 5 Aug 2026 16:51:59 +0800 Subject: [PATCH 05/15] feat(swagger): generate localized static API schema --- apps/edu/views.py | 6 + apps/system/views.py | 1 + apps/third/views.py | 3 + apps/utils/management/__init__.py | 0 apps/utils/management/commands/__init__.py | 0 .../management/commands/build_swagger.py | 71 ++++ apps/utils/swagger.py | 367 ++++++++++++++++++ apps/utils/test_swagger.py | 142 +++++++ apps/utils/viewsets.py | 5 +- server/settings.py | 20 + server/swagger.py | 12 + server/urls.py | 15 +- 12 files changed, 632 insertions(+), 10 deletions(-) create mode 100644 apps/utils/management/__init__.py create mode 100644 apps/utils/management/commands/__init__.py create mode 100644 apps/utils/management/commands/build_swagger.py create mode 100644 apps/utils/swagger.py create mode 100644 apps/utils/test_swagger.py create mode 100644 server/swagger.py diff --git a/apps/edu/views.py b/apps/edu/views.py index c6fa69ef..00da4c56 100644 --- a/apps/edu/views.py +++ b/apps/edu/views.py @@ -71,6 +71,8 @@ class ExamViewSet(CustomModelViewSet): def get_queryset(self): qs = super().get_queryset() + if getattr(self, 'swagger_fake_view', False): + return qs if has_perm(self.request.user, ["exam.view"]): return qs user:User = self.request.user @@ -142,6 +144,8 @@ class ExamRecordViewSet(ListModelMixin, DestroyModelMixin, RetrieveModelMixin, C def get_queryset(self): qs = super().get_queryset() + if getattr(self, 'swagger_fake_view', False): + return qs if has_perm(self.request.user, ["examrecord.view"]): return qs return qs.filter(create_by=self.request.user) @@ -207,6 +211,8 @@ class TrainRecordViewSet(CustomModelViewSet): def get_queryset(self): qs = super().get_queryset() + if getattr(self, 'swagger_fake_view', False): + return qs if has_perm(self.request.user, ["train.view"]): return qs return qs.filter(create_by=self.request.user) diff --git a/apps/system/views.py b/apps/system/views.py index 32742d88..c9783a0d 100755 --- a/apps/system/views.py +++ b/apps/system/views.py @@ -651,6 +651,7 @@ class FileViewSet(BulkCreateModelMixin, RetrieveModelMixin, CustomListModelMixin class ApkViewSet(MyLoggingMixin, CustomListModelMixin, BulkCreateModelMixin, GenericViewSet): perms_map = {'get': '*', 'post': 'apk.upload'} serializer_class = ApkSerializer + filter_backends = [] def get_authenticators(self): if self.request.method == 'GET': diff --git a/apps/third/views.py b/apps/third/views.py index 6346e094..3aca1206 100755 --- a/apps/third/views.py +++ b/apps/third/views.py @@ -69,6 +69,7 @@ class SpeakerViewSet(CustomGenericViewSet): """ perms_map = {} serializer_class = serializers.Serializer + filter_backends = [] @action(methods=['get'], detail=False, permission_classes=[IsAuthenticated]) @@ -125,6 +126,7 @@ class XxTestView(APIView): class XxCommonViewSet(CreateModelMixin, CustomGenericViewSet): perms_map = {'post': '*'} serializer_class = RequestCommonSerializer + filter_backends = [] def create(self, request, *args, **kwargs): """ @@ -258,6 +260,7 @@ class KingCommonViewSet(CreateModelMixin, CustomGenericViewSet): class DhCommonViewSet(CreateModelMixin, CustomGenericViewSet): perms_map = {'post': '*'} serializer_class = RequestCommonSerializer + filter_backends = [] def create(self, request, *args, **kwargs): """ diff --git a/apps/utils/management/__init__.py b/apps/utils/management/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/apps/utils/management/commands/__init__.py b/apps/utils/management/commands/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/apps/utils/management/commands/build_swagger.py b/apps/utils/management/commands/build_swagger.py new file mode 100644 index 00000000..c47fcf98 --- /dev/null +++ b/apps/utils/management/commands/build_swagger.py @@ -0,0 +1,71 @@ +import json +import os +from io import StringIO +from pathlib import Path + +from django.conf import settings +from django.core.management import BaseCommand, CommandError, call_command + + +class Command(BaseCommand): + help = "生成供 Swagger UI 和 ReDoc 使用的静态 Swagger JSON" + + def add_arguments(self, parser): + parser.add_argument( + "--output", + help="输出路径,默认使用 settings.SWAGGER_SCHEMA_PATH", + ) + parser.add_argument( + "--url", + help="文档中的 API 根地址,默认使用 settings.BASE_URL", + ) + + def handle(self, *args, **options): + target = Path(options["output"] or settings.SWAGGER_SCHEMA_PATH) + if not target.is_absolute(): + target = Path(settings.BASE_DIR) / target + target = target.resolve() + target.parent.mkdir(parents=True, exist_ok=True) + + temporary = target.with_name(f".{target.name}.{os.getpid()}.tmp") + try: + output = StringIO() + call_command( + "generate_swagger", + "-", + format="json", + api_url=options["url"] or settings.BASE_URL, + mock=True, + verbosity=0, + stdout=output, + ) + content = output.getvalue() + schema = json.loads(content) + if schema.get("swagger") != "2.0" or not schema.get("paths"): + raise CommandError("生成的 Swagger 文档缺少版本或接口路径") + + content = json.dumps( + schema, + ensure_ascii=False, + separators=(",", ":"), + ) + temporary.write_text(content, encoding="utf-8") + os.replace(temporary, target) + except Exception as exc: + if isinstance(exc, CommandError): + raise + raise CommandError(f"生成 Swagger 文档失败:{exc}") from exc + finally: + temporary.unlink(missing_ok=True) + + operation_count = sum( + method.lower() in {"get", "post", "put", "patch", "delete"} + for path in schema["paths"].values() + for method in path + ) + self.stdout.write( + self.style.SUCCESS( + f"Swagger文档已生成:{target} " + f"({len(schema['paths'])}个路径,{operation_count}个操作)" + ) + ) diff --git a/apps/utils/swagger.py b/apps/utils/swagger.py new file mode 100644 index 00000000..59448844 --- /dev/null +++ b/apps/utils/swagger.py @@ -0,0 +1,367 @@ +import re +from pathlib import Path + +from django.apps import apps +from django.conf import settings +from django.http import FileResponse, JsonResponse +from drf_yasg import openapi +from drf_yasg.inspectors import FieldInspector, SwaggerAutoSchema +from drf_yasg.inspectors.base import NotHandled + + +CRUD_SUMMARIES = { + "list": "查询{resource}列表", + "retrieve": "查询{resource}详情", + "create": "新增{resource}", + "update": "更新{resource}", + "partial_update": "部分更新{resource}", + "destroy": "删除{resource}", +} + +TAG_NAMES = { + "am": "区域与准入管理", + "asm": "资产管理", + "cm": "标签管理", + "cms": "内容管理", + "develop": "开发工具", + "ecm": "事件管理", + "edu": "培训考试", + "em": "设备管理", + "enm": "能源管理", + "inm": "库存管理", + "mpr": "物资申购与领用", + "mtm": "物料与工艺管理", + "ofm": "办公管理", + "opm": "作业许可", + "pm": "生产任务管理", + "pum": "采购管理", + "qm": "质量管理", + "rem": "研发项目管理", + "rpm": "相关方管理", + "sam": "销售管理", + "third": "第三方集成", + "utils": "通用工具", + "wpm": "生产管理", + "wpmw": "动态产品管理", + "file": "文件管理", +} + +FIELD_NAMES = { + "id": "主键ID", + "ids": "主键ID列表", + "access": "访问令牌", + "refresh": "刷新令牌", + "password_check": "密码确认", + "base64": "Base64数据", + "server_time": "服务器时间", + "timezone": "时区", + "next": "下一页", + "previous": "上一页", + "results": "结果列表", + "detail": "详情", + "items": "明细列表", + "files": "附件列表", + "echart_options": "图表配置", + "tdata_list": "数据列表", + "page": "页码", + "page_size": "每页数量", + "ordering": "排序字段", + "querys": "查询条件列表", + "annotate_field_list": "聚合字段列表", +} + +FIELD_TOKENS = { + "name": "名称", + "code": "编码", + "number": "编号", + "description": "说明", + "note": "备注", + "employee": "人员", + "user": "用户", + "leader": "负责人", + "manager": "负责人", + "keeper": "保管人", + "participant": "参与人", + "post": "岗位", + "dept": "部门", + "belong": "所属", + "create": "创建", + "update": "更新", + "submit": "提交", + "handle": "处理", + "test": "检验", + "material": "物料", + "supplier": "供应商", + "defect": "缺陷", + "equipment": "设备", + "warehouse": "仓库", + "process": "工序", + "operation": "操作", + "state": "状态", + "cate": "分类", + "type": "类型", + "area": "区域", + "team": "班组", + "shift": "班次", + "ticket": "工单", + "file": "文件", + "photo": "照片", + "image": "图片", + "origin": "来源", + "in": "入库", + "out": "出库", + "list": "列表", + "count": "数量", + "total": "总计", + "enabled": "是否启用", +} + +QUERY_PARAMETERS = { + "page": "页码", + "page_size": "每页数量", + "search": "搜索关键字", + "ordering": "排序字段,字段名前加“-”表示倒序", + "format": "响应格式", +} + +LOOKUP_NAMES = { + "in": "属于列表", + "contains": "包含", + "icontains": "包含(忽略大小写)", + "gte": "大于或等于", + "gt": "大于", + "lte": "小于或等于", + "lt": "小于", + "isnull": "是否为空", + "exact": "等于", +} + + +def _contains_chinese(value): + return bool(re.search(r"[\u4e00-\u9fff]", str(value or ""))) + + +def swagger_schema_file(request): + schema_path = Path(settings.SWAGGER_SCHEMA_PATH) + if not schema_path.is_file(): + return JsonResponse( + {"detail": "Swagger文档尚未生成,请先运行 manage.py build_swagger"}, + status=503, + ) + + response = FileResponse( + schema_path.open("rb"), + content_type="application/json; charset=utf-8", + filename="swagger.json", + ) + response["Content-Disposition"] = 'inline; filename="swagger.json"' + response["Cache-Control"] = "no-cache" + return response + + +def _serializer_model(field): + parent = getattr(field, "parent", None) + while parent is not None: + meta = getattr(parent, "Meta", None) + model = getattr(meta, "model", None) + if model is not None: + return model + parent = getattr(parent, "parent", None) + return None + + +def _model_path_label(model, parts): + labels = [] + for part in parts: + if model is None: + break + try: + model_field = model._meta.get_field(part) + except Exception: + break + verbose_name = getattr(model_field, "verbose_name", "") + if _contains_chinese(verbose_name): + labels.append(str(verbose_name)) + model = getattr(model_field, "related_model", None) + return " / ".join(labels) + + +def _field_name_label(field_name): + field_name = str(field_name or "").strip("_") + if field_name in FIELD_NAMES: + return FIELD_NAMES[field_name] + tokens = field_name.split("_") + if tokens and all(token in FIELD_TOKENS for token in tokens): + return "".join(FIELD_TOKENS[token] for token in tokens) + return "" + + +class ChineseFieldInspector(FieldInspector): + """优先使用模型字段中文名称补全 serializer 字段标题。""" + + def field_to_swagger_object(self, field, **kwargs): + return NotHandled + + def process_result(self, result, method_name, obj, **kwargs): + if ( + method_name != "field_to_swagger_object" + or not isinstance(result, openapi.SwaggerDict) + or "$ref" in result + or _contains_chinese(result.get("title")) + ): + return result + + source_attrs = getattr(obj, "source_attrs", None) or [] + model_label = _model_path_label(_serializer_model(obj), source_attrs) + label = model_label or _field_name_label(getattr(obj, "field_name", "")) + field_name = getattr(obj, "field_name", "") + if label: + result["title"] = label + elif field_name: + result["title"] = f"{field_name}(字段)" + return result + + +class ChineseSwaggerAutoSchema(SwaggerAutoSchema): + """为未显式编写文档的接口补充稳定、可读的中文展示信息。""" + + field_inspectors = [ChineseFieldInspector] + SwaggerAutoSchema.field_inspectors + + def get_operation(self, operation_keys=None): + operation = super().get_operation(operation_keys) + model = getattr(getattr(self.view, "queryset", None), "model", None) + for parameter in operation.get("parameters", []): + current = parameter.get("description", "") + if _contains_chinese(current): + continue + location = parameter.get("in") + if location == openapi.IN_BODY: + description = "请求数据" + elif location == openapi.IN_PATH: + description = f"路径参数:{parameter.get('name', '')}" + else: + description = self._get_parameter_description( + parameter.get("name", ""), model + ) + if current: + description = f"{description};{current}" + parameter["description"] = description + return operation + + def get_summary_and_description(self): + summary, description = super().get_summary_and_description() + if summary: + if description and not _contains_chinese(description): + description = f"{summary}\n\n{description}" + return summary, description or summary + + resource = self._get_resource_name() + action = getattr(self.view, "action", None) + template = CRUD_SUMMARIES.get(action) + if template: + summary = template.format(resource=resource) + elif resource: + action_name = str(action or self.method).replace("_", " ") + display_resource = resource + if not _contains_chinese(display_resource): + display_resource = f"{display_resource}接口" + summary = f"{display_resource}:{action_name}" + + if description and not _contains_chinese(description): + description = f"{summary}\n\n{description}" + return summary, description or summary + + def get_request_body_parameters(self, consumes): + parameters = super().get_request_body_parameters(consumes) + for parameter in parameters: + if parameter.get("in") == openapi.IN_BODY and not _contains_chinese( + parameter.get("description") + ): + parameter["description"] = "请求数据" + return parameters + + def get_query_parameters(self): + parameters = super().get_query_parameters() + model = getattr(getattr(self.view, "queryset", None), "model", None) + for parameter in parameters: + current = parameter.get("description", "") + if _contains_chinese(current): + continue + description = self._get_parameter_description( + parameter.get("name", ""), model + ) + if current: + description = f"{description};{current}" + parameter["description"] = description + return parameters + + def get_responses(self): + responses = super().get_responses() + descriptions = { + "200": "请求成功", + "201": "创建成功", + "202": "请求已接受", + "204": "操作成功,无响应内容", + "400": "请求参数错误", + "401": "未认证或认证已失效", + "403": "无权访问", + "404": "资源不存在", + } + for status, response in responses.items(): + if not response.get("description"): + response["description"] = descriptions.get(str(status), "接口响应") + return responses + + def _get_parameter_description(self, name, model): + if name in QUERY_PARAMETERS: + return QUERY_PARAMETERS[name] + + parts = str(name).split("__") + lookup = LOOKUP_NAMES.get(parts[-1]) + field_parts = parts[:-1] if lookup else parts + label = _model_path_label(model, field_parts) + if not label: + label = _field_name_label(field_parts[-1] if field_parts else name) + if not label: + label = f"查询参数:{name}" + if lookup: + label = f"{label}({lookup})" + return label + + def get_tags(self, operation_keys=None): + tags = super().get_tags(operation_keys) + if self.overrides.get("tags") or not tags: + return tags + + if tags[0] in TAG_NAMES: + return [TAG_NAMES[tags[0]]] + + try: + app_config = apps.get_app_config(tags[0]) + except LookupError: + return tags + + if _contains_chinese(app_config.verbose_name): + return [str(app_config.verbose_name)] + return tags + + def _get_resource_name(self): + queryset = getattr(self.view, "queryset", None) + model = getattr(queryset, "model", None) + + if model is None: + serializer_class = getattr(self.view, "serializer_class", None) + meta = getattr(serializer_class, "Meta", None) + model = getattr(meta, "model", None) + + if model is None: + return "接口" + + match = re.search(r"TN\s*[::]\s*([^\n\r]+)", model.__doc__ or "") + if match: + return match.group(1).strip() + + verbose_name = str(model._meta.verbose_name) + if _contains_chinese(verbose_name): + return verbose_name + return model.__name__ diff --git a/apps/utils/test_swagger.py b/apps/utils/test_swagger.py new file mode 100644 index 00000000..27afaeaf --- /dev/null +++ b/apps/utils/test_swagger.py @@ -0,0 +1,142 @@ +import json +from pathlib import Path +from tempfile import TemporaryDirectory +from types import SimpleNamespace +from unittest.mock import patch + +from django.conf import settings +from django.core.management import call_command +from django.test import SimpleTestCase, override_settings + +from apps.am.models import Area +from apps.am.views import AreaViewSet +from apps.utils.swagger import ChineseSwaggerAutoSchema, swagger_schema_file + + +class ChineseSwaggerAutoSchemaTests(SimpleTestCase): + def make_schema(self, view, method="GET"): + schema = ChineseSwaggerAutoSchema.__new__(ChineseSwaggerAutoSchema) + schema.view = view + schema.method = method + schema.path = "/am/area/" + schema.overrides = {} + schema.operation_keys = ("am", "area", "list") + schema._sch = SimpleNamespace(get_description=lambda path, method: "") + return schema + + def test_crud_summary_uses_model_chinese_name(self): + view = SimpleNamespace(queryset=Area.objects.all(), action="list") + schema = self.make_schema(view) + + summary, description = schema.get_summary_and_description() + + self.assertEqual(summary, "查询地图区域列表") + self.assertEqual(description, "查询地图区域列表") + + def test_explicit_summary_takes_priority(self): + view = SimpleNamespace(queryset=Area.objects.all(), action="list") + schema = self.make_schema(view) + schema.overrides = { + "operation_summary": "区域自定义查询", + "operation_description": "自定义说明", + } + + summary, description = schema.get_summary_and_description() + + self.assertEqual(summary, "区域自定义查询") + self.assertEqual(description, "自定义说明") + + def test_custom_action_with_english_model_name_has_chinese_hint(self): + model = SimpleNamespace( + __doc__="", + __name__="Dataset", + _meta=SimpleNamespace(verbose_name="dataset"), + ) + queryset = SimpleNamespace(model=model) + view = SimpleNamespace(queryset=queryset, action="base") + schema = self.make_schema(view) + + summary, _ = schema.get_summary_and_description() + + self.assertEqual(summary, "Dataset接口:base") + + def test_tag_uses_chinese_business_module_name(self): + view = SimpleNamespace(queryset=Area.objects.all(), action="list") + schema = self.make_schema(view) + + self.assertEqual(schema.get_tags(("am", "area", "list")), ["区域与准入管理"]) + + def test_filter_parameter_uses_model_field_labels(self): + view = SimpleNamespace(queryset=Area.objects.all(), action="list") + schema = self.make_schema(view) + + description = schema._get_parameter_description( + "manager__name__contains", + Area, + ) + + self.assertIn("区域负责人", description) + self.assertIn("包含", description) + + def test_swagger_queryset_skips_permission_data_lookup(self): + view = AreaViewSet(basename="area") + view.action = "list" + view.swagger_fake_view = True + + with patch("apps.utils.viewsets.get_user_perms_map") as permission_lookup: + queryset = view.get_queryset() + + self.assertIs(queryset.model, Area) + permission_lookup.assert_not_called() + + +class SwaggerSettingsTests(SimpleTestCase): + def test_swagger_supports_jwt_authorization_header(self): + from django.conf import settings + + bearer = settings.SWAGGER_SETTINGS["SECURITY_DEFINITIONS"]["Bearer"] + + self.assertEqual(bearer["type"], "apiKey") + self.assertEqual(bearer["name"], "Authorization") + self.assertEqual(bearer["in"], "header") + + def test_swagger_ui_uses_static_schema(self): + from django.conf import settings + + self.assertEqual(settings.SWAGGER_SETTINGS["SPEC_URL"], "schema-swagger-json") + self.assertEqual(settings.REDOC_SETTINGS["SPEC_URL"], "schema-swagger-json") + + +class BuildSwaggerCommandTests(SimpleTestCase): + def test_command_writes_valid_utf8_schema(self): + schema = { + "swagger": "2.0", + "info": {"title": "中文文档"}, + "paths": {"/demo/": {"get": {}}}, + } + + def generate_schema(command_name, output_file, **options): + self.assertEqual(command_name, "generate_swagger") + self.assertEqual(output_file, "-") + options["stdout"].write(json.dumps(schema, ensure_ascii=False)) + + with TemporaryDirectory(dir=settings.BASE_DIR) as directory: + target = Path(directory) / "swagger.json" + with override_settings(SWAGGER_SCHEMA_PATH=str(target)): + with patch( + "apps.utils.management.commands.build_swagger.call_command", + side_effect=generate_schema, + ): + call_command("build_swagger", verbosity=0) + + content = target.read_text(encoding="utf-8") + self.assertIn("中文文档", content) + self.assertEqual(json.loads(content), schema) + + with override_settings(SWAGGER_SCHEMA_PATH=str(target)): + response = swagger_schema_file(SimpleNamespace()) + body = b"".join(response.streaming_content) + response.close() + + self.assertEqual(response.status_code, 200) + self.assertEqual(json.loads(body), schema) diff --git a/apps/utils/viewsets.py b/apps/utils/viewsets.py index a4246422..93623816 100755 --- a/apps/utils/viewsets.py +++ b/apps/utils/viewsets.py @@ -154,6 +154,9 @@ class CustomGenericViewSet(MyLoggingMixin, GenericViewSet): def get_queryset(self): queryset = super().get_queryset() queryset = self.get_queryset_custom(queryset) + # drf-yasg 生成文档时不应读取权限或业务数据。 + if getattr(self, 'swagger_fake_view', False): + return queryset if self.data_filter: user = self.request.user if user.is_superuser: @@ -232,4 +235,4 @@ class EuModelViewSet(BulkCreateModelMixin, CustomListModelMixin, CustomRetrieveModelMixin, BulkDestroyModelMixin, ComplexQueryMixin, CustomGenericViewSet): """ 不支持更新的增强ModelViewSet - """ \ No newline at end of file + """ diff --git a/server/settings.py b/server/settings.py index d0bf321c..1df8238b 100755 --- a/server/settings.py +++ b/server/settings.py @@ -178,6 +178,7 @@ USE_TZ = True STATIC_URL = '/static/' STATIC_ROOT = os.path.join(BASE_DIR, 'dist/static') +SWAGGER_SCHEMA_PATH = os.path.join(STATIC_ROOT, 'openapi/swagger.json') # STATICFILES_DIRS = ( # os.path.join(BASE_DIR, 'dist/static'), # ) @@ -267,8 +268,27 @@ CELERYD_SOFT_TIME_LIMIT = 60*10 # swagger配置 SWAGGER_SETTINGS = { + 'DEFAULT_INFO': 'server.swagger.api_info', + 'DEFAULT_API_URL': BASE_URL, + 'SPEC_URL': 'schema-swagger-json', 'LOGIN_URL': '/django/admin/login/', 'LOGOUT_URL': '/django/admin/logout/', + 'DEFAULT_AUTO_SCHEMA_CLASS': 'apps.utils.swagger.ChineseSwaggerAutoSchema', + 'SECURITY_DEFINITIONS': { + 'Bearer': { + 'type': 'apiKey', + 'name': 'Authorization', + 'in': 'header', + 'description': 'JWT认证,请输入:Bearer ', + }, + 'Basic': { + 'type': 'basic', + }, + }, +} + +REDOC_SETTINGS = { + 'SPEC_URL': 'schema-swagger-json', } # 日志配置 diff --git a/server/swagger.py b/server/swagger.py new file mode 100644 index 00000000..c1c20132 --- /dev/null +++ b/server/swagger.py @@ -0,0 +1,12 @@ +from django.conf import settings +from drf_yasg import openapi + +from server.settings import get_sysconfig + + +api_info = openapi.Info( + title=f'{settings.SYS_NAME}--{get_sysconfig("base.base_name", "demo")}', + default_version=settings.SYS_VERSION, + contact=openapi.Contact(email="caoqianming@foxmail.com"), + license=openapi.License(name="MIT License"), +) diff --git a/server/urls.py b/server/urls.py index ac46779f..b4fab73e 100755 --- a/server/urls.py +++ b/server/urls.py @@ -17,19 +17,14 @@ from django.conf import settings from django.conf.urls.static import static from django.contrib import admin from django.urls import include, path -from drf_yasg import openapi from drf_yasg.views import get_schema_view from rest_framework.documentation import include_docs_urls from django.views.generic import TemplateView -from server.settings import get_sysconfig +from apps.utils.swagger import swagger_schema_file +from server.swagger import api_info schema_view = get_schema_view( - openapi.Info( - title=f'{settings.SYS_NAME}--{get_sysconfig("base.base_name", "demo")}', - default_version=settings.SYS_VERSION, - contact=openapi.Contact(email="caoqianming@foxmail.com"), - license=openapi.License(name="MIT License"), - ), + api_info, public=True, permission_classes=[], url=settings.BASE_URL @@ -89,8 +84,10 @@ urlpatterns = [ if getattr(settings, 'ENABLE_SWAGGER', True): urlpatterns += [ # api文档 + path('api/swagger.json', swagger_schema_file, + name='schema-swagger-json'), path('api/swagger/', schema_view.with_ui('swagger', cache_timeout=0), name='schema-swagger-ui'), path('api/redoc/', schema_view.with_ui('redoc', cache_timeout=0), name='schema-redoc'), - ] \ No newline at end of file + ] From b365e063183b8bacbb8e5efc6861f3d27e58ac64 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Thu, 6 Aug 2026 10:31:37 +0800 Subject: [PATCH 06/15] feat(inventory): add effective defect grade filtering --- apps/inm/filters.py | 6 ++ apps/inm/serializers.py | 20 ++++++- apps/inm/tests.py | 70 ++++++++++++++++++++++- apps/inm/views.py | 2 +- apps/qm/defect_grades.py | 42 ++++++++++++++ apps/qm/models.py | 16 ++++-- apps/wpm/filters.py | 5 ++ apps/wpm/serializers.py | 5 +- apps/wpm/tests.py | 51 ++++++++++++++++- apps/wpm/tests/test_defect_grade.py | 86 +++++++++++++++++++++++++++++ 10 files changed, 290 insertions(+), 13 deletions(-) create mode 100644 apps/qm/defect_grades.py create mode 100644 apps/wpm/tests/test_defect_grade.py diff --git a/apps/inm/filters.py b/apps/inm/filters.py index a7455474..a0b9e092 100644 --- a/apps/inm/filters.py +++ b/apps/inm/filters.py @@ -1,10 +1,16 @@ from django_filters import rest_framework as filters from apps.inm.models import MaterialBatch, MIO from django.db.models import Q, Subquery, OuterRef, F +from apps.qm.defect_grades import effective_defect_grade_q class MaterialBatchFilter(filters.FilterSet): count_canmio__gt = filters.NumberFilter( method='filter_count_canmio__gt', label='可发数量大于') + defect_grade = filters.NumberFilter( + method='filter_defect_grade', label='有效缺陷等级') + + def filter_defect_grade(self, queryset, name, value): + return queryset.filter(effective_defect_grade_q(value)) class Meta: model = MaterialBatch diff --git a/apps/inm/serializers.py b/apps/inm/serializers.py index 018673eb..326c2fb7 100644 --- a/apps/inm/serializers.py +++ b/apps/inm/serializers.py @@ -14,6 +14,7 @@ from django.db.models import F, Sum, DecimalField from server.settings import get_sysconfig from apps.wpmw.models import Wpr from decimal import Decimal +from apps.qm.defect_grades import DEFECT_GRADE_NAMES, effective_defect_grade class WareHourseSerializer(CustomModelSerializer): @@ -49,6 +50,8 @@ class MaterialBatchSerializer(CustomModelSerializer): source='supplier', read_only=True) material_ = MaterialSerializer(source='material', read_only=True) defect_name = serializers.CharField(source="defect.name", read_only=True) + defect_grade = serializers.SerializerMethodField() + defect_grade_name = serializers.SerializerMethodField() count_mioing = serializers.SerializerMethodField(label='正在出入库数量') class Meta: @@ -61,6 +64,12 @@ class MaterialBatchSerializer(CustomModelSerializer): # 保留 decimal 精度(原 IntegerField 会截断在途量, 导致可发量偏大) return instance.count_mioing_anno if hasattr(instance, 'count_mioing_anno') else instance.count_mioing + def get_defect_grade(self, instance): + return effective_defect_grade(instance) + + def get_defect_grade_name(self, instance): + return DEFECT_GRADE_NAMES[self.get_defect_grade(instance)] + def to_representation(self, instance): ret = super().to_representation(instance) if 'count' in ret: @@ -86,6 +95,15 @@ class MaterialBatchDetailSerializer(CustomModelSerializer): source='a_mb', read_only=True, many=True) supplier_name = serializers.StringRelatedField( source='supplier', read_only=True) + defect_name = serializers.CharField(source="defect.name", read_only=True) + defect_grade = serializers.SerializerMethodField() + defect_grade_name = serializers.SerializerMethodField() + + def get_defect_grade(self, instance): + return effective_defect_grade(instance) + + def get_defect_grade_name(self, instance): + return DEFECT_GRADE_NAMES[self.get_defect_grade(instance)] class Meta: model = MaterialBatch @@ -542,4 +560,4 @@ class PackSerializer(CustomModelSerializer): class PackMioSerializer(serializers.Serializer): mioitems = serializers.ListField(child=serializers.CharField(), label="明细ID") pack_index = serializers.IntegerField(label="包装箱序号") - # pack = serializers.CharField(label="包装箱ID") \ No newline at end of file + # pack = serializers.CharField(label="包装箱ID") diff --git a/apps/inm/tests.py b/apps/inm/tests.py index c00fa963..c4c5b2f4 100644 --- a/apps/inm/tests.py +++ b/apps/inm/tests.py @@ -4,10 +4,78 @@ from threading import Barrier from unittest import skipUnless from django.db import connection, connections, transaction -from django.test import SimpleTestCase, TransactionTestCase +from django.test import SimpleTestCase, TestCase, TransactionTestCase +from apps.inm.filters import MaterialBatchFilter from apps.inm.models import MaterialBatch, WareHouse +from apps.inm.serializers import MaterialBatchSerializer from apps.mtm.models import Material +from apps.qm.models import Defect + + +class MaterialBatchDefectGradeTests(TestCase): + @classmethod + def setUpTestData(cls): + cls.material = Material.objects.create(name='仓库缺陷等级测试物料') + cls.warehouse = WareHouse.objects.create( + number='GRADE', + name='等级测试仓库', + place='测试地点', + ) + cls.defect_b = Defect.objects.create( + name='仓库B类缺陷', + cate=Defect.cate_list[0], + okcate=Defect.DEFECT_OK_B, + ) + cls.notok_without_defect = MaterialBatch.objects.create( + material=cls.material, + warehouse=cls.warehouse, + batch='MB-NOTOK-NONE', + count=1, + state=20, + ) + cls.normal_with_b_defect = MaterialBatch.objects.create( + material=cls.material, + warehouse=cls.warehouse, + batch='MB-NORMAL-B', + count=1, + state=10, + defect=cls.defect_b, + ) + + def test_serializer_uses_defect_or_defaults_to_ok_independent_of_state(self): + no_defect_data = MaterialBatchSerializer( + self.notok_without_defect + ).data + b_defect_data = MaterialBatchSerializer( + self.normal_with_b_defect + ).data + + self.assertEqual(no_defect_data['defect_grade'], Defect.DEFECT_OK) + self.assertEqual(no_defect_data['defect_grade_name'], '合格') + self.assertEqual(b_defect_data['defect_grade'], Defect.DEFECT_OK_B) + self.assertEqual(b_defect_data['defect_grade_name'], '合格B类') + + def test_effective_grade_filter_is_independent_of_state(self): + ok_items = MaterialBatchFilter( + {'defect_grade': Defect.DEFECT_OK}, + queryset=MaterialBatch.objects.all(), + ).qs + b_items = MaterialBatchFilter( + {'defect_grade': Defect.DEFECT_OK_B}, + queryset=MaterialBatch.objects.all(), + ).qs + + self.assertQuerySetEqual( + ok_items, + [self.notok_without_defect], + transform=lambda item: item, + ) + self.assertQuerySetEqual( + b_items, + [self.normal_with_b_defect], + transform=lambda item: item, + ) class MaterialBatchInventoryKeyTests(SimpleTestCase): diff --git a/apps/inm/views.py b/apps/inm/views.py index c6984703..ef17d831 100644 --- a/apps/inm/views.py +++ b/apps/inm/views.py @@ -60,7 +60,7 @@ class MaterialBatchViewSet(ListModelMixin, CustomGenericViewSet): queryset = MaterialBatch.objects.filter(count__gt=0) serializer_class = MaterialBatchSerializer retrieve_serializer_class = MaterialBatchDetailSerializer - select_related_fields = ['warehouse', 'material', 'supplier'] + select_related_fields = ['warehouse', 'material', 'supplier', 'defect'] filterset_class = MaterialBatchFilter search_fields = ['material__name', 'material__number', 'material__model', 'material__specification', 'batch'] diff --git a/apps/qm/defect_grades.py b/apps/qm/defect_grades.py new file mode 100644 index 00000000..665d9bb7 --- /dev/null +++ b/apps/qm/defect_grades.py @@ -0,0 +1,42 @@ +from django.db.models import Q + + +DEFECT_OK = 10 +DEFECT_OK_B = 20 +DEFECT_NOTOK = 30 + +DEFECT_GRADE_CHOICES = ( + (DEFECT_OK, "合格"), + (DEFECT_OK_B, "合格B类"), + (DEFECT_NOTOK, "不合格"), +) +DEFECT_GRADE_NAMES = dict(DEFECT_GRADE_CHOICES) + + +def effective_defect_grade(instance, notok_sign_field=None): + """Return the inventory grade without coupling it to inventory state.""" + defect = getattr(instance, "defect", None) + if defect is not None: + return defect.okcate + if notok_sign_field and getattr(instance, notok_sign_field, None): + return DEFECT_NOTOK + return DEFECT_OK + + +def effective_defect_grade_q(value, notok_sign_field=None): + """Build an index-friendly query matching ``effective_defect_grade``.""" + explicit_grade = Q(defect__okcate=value) + without_defect = Q(defect__isnull=True) + + if not notok_sign_field: + return explicit_grade | without_defect if value == DEFECT_OK else explicit_grade + + has_legacy_sign = ( + Q(**{f"{notok_sign_field}__isnull": False}) + & ~Q(**{notok_sign_field: ""}) + ) + if value == DEFECT_OK: + return explicit_grade | (without_defect & ~has_legacy_sign) + if value == DEFECT_NOTOK: + return explicit_grade | (without_defect & has_legacy_sign) + return explicit_grade diff --git a/apps/qm/models.py b/apps/qm/models.py index 49f5bfe8..e7a51ec0 100644 --- a/apps/qm/models.py +++ b/apps/qm/models.py @@ -8,19 +8,25 @@ from django.utils.translation import gettext_lazy as _ from django.db import transaction from django.db.models import Sum from rest_framework.exceptions import ParseError +from apps.qm.defect_grades import ( + DEFECT_GRADE_CHOICES, + DEFECT_NOTOK as GRADE_NOTOK, + DEFECT_OK as GRADE_OK, + DEFECT_OK_B as GRADE_OK_B, +) class Defect(CommonAModel): """TN:缺陷项""" - DEFECT_OK = 10 - DEFECT_OK_B = 20 - DEFECT_NOTOK = 30 + DEFECT_OK = GRADE_OK + DEFECT_OK_B = GRADE_OK_B + DEFECT_NOTOK = GRADE_NOTOK cate_list = ["尺寸", "外观", "内质", "性能"] name = models.CharField(max_length=50, verbose_name="名称") code = models.CharField(max_length=50, verbose_name="标识", null=True, blank=True) cate = models.CharField(max_length=50, verbose_name="分类", help_text=str(cate_list)) okcate= models.PositiveSmallIntegerField(verbose_name="不合格分类", - choices=((DEFECT_OK, "合格"), (DEFECT_OK_B, "合格B类"), (DEFECT_NOTOK, "不合格")), - default=DEFECT_NOTOK) + choices=DEFECT_GRADE_CHOICES, + default=GRADE_NOTOK) note = models.TextField('备注', null=True, blank=True) def __str__(self): diff --git a/apps/wpm/filters.py b/apps/wpm/filters.py index 705ef4f6..50a6cafd 100644 --- a/apps/wpm/filters.py +++ b/apps/wpm/filters.py @@ -5,6 +5,7 @@ from apps.mtm.models import Route, Material from django.db.models import Q, Exists, OuterRef from rest_framework.exceptions import ParseError from datetime import datetime +from apps.qm.defect_grades import effective_defect_grade_q class SfLogFilter(filters.FilterSet): class Meta: @@ -44,6 +45,10 @@ class WMaterialFilter(filters.FilterSet): mlog_date_start = filters.DateFilter(label="产出开始", method="filter_mlog_date_start") mlog_date_end = filters.DateFilter(label="产出结束", method="filter_mlog_date_end") current_merged = filters.BooleanFilter(label="是否本工段新合成的批", method="filter_current_merged") + defect_grade = filters.NumberFilter(label="有效缺陷等级", method="filter_defect_grade") + + def filter_defect_grade(self, queryset, name, value): + return queryset.filter(effective_defect_grade_q(value, "notok_sign")) def filter_mlog_date_start(self, queryset, name, value): mgroupId = self.data.get("mgroup", None) diff --git a/apps/wpm/serializers.py b/apps/wpm/serializers.py index 53d4daad..bd868f85 100644 --- a/apps/wpm/serializers.py +++ b/apps/wpm/serializers.py @@ -24,6 +24,7 @@ from apps.wpmw.models import Wpr from apps.qm.serializers import FtestProcessSerializer, FtestProcessListSerializer import logging from apps.qm.models import Defect +from apps.qm.defect_grades import DEFECT_GRADE_NAMES, effective_defect_grade from apps.utils.snowflake import idWorker from decimal import Decimal from apps.em.models import Equipment @@ -199,10 +200,10 @@ class WMaterialSerializer(CustomModelSerializer): return getattr(NotOkOption, obj.notok_sign, NotOkOption.qt).label if obj.notok_sign else None def get_defect_grade(self, obj): - return obj.defect.okcate if obj.defect else None + return effective_defect_grade(obj, "notok_sign") def get_defect_grade_name(self, obj): - return obj.defect.get_okcate_display() if obj.defect else None + return DEFECT_GRADE_NAMES[self.get_defect_grade(obj)] def get_count_working(self, obj): # 列表接口 queryset 已注解(单次聚合); 嵌套等无注解场景回退模型属性 diff --git a/apps/wpm/tests.py b/apps/wpm/tests.py index 6954cd04..9ffaf96f 100644 --- a/apps/wpm/tests.py +++ b/apps/wpm/tests.py @@ -138,13 +138,30 @@ class WMaterialDefectGradeTests(TestCase): count=1, state=WMaterial.WM_OK, ) + cls.repair_without_defect = WMaterial.objects.create( + material=cls.material, + batch="REPAIR-NONE", + count=1, + state=WMaterial.WM_REPAIR, + ) + cls.notok_with_legacy_sign = WMaterial.objects.create( + material=cls.material, + batch="NOTOK-LEGACY", + count=1, + state=WMaterial.WM_NOTOK, + notok_sign="zw", + ) - def test_serializer_exposes_nullable_defect_grade_without_using_state(self): + def test_serializer_exposes_effective_defect_grade_without_using_state(self): normal_notok_data = WMaterialSerializer( self.normal_with_notok_defect ).data notok_b_data = WMaterialSerializer(self.notok_with_b_defect).data no_defect_data = WMaterialSerializer(self.normal_without_defect).data + repair_no_defect_data = WMaterialSerializer( + self.repair_without_defect + ).data + legacy_data = WMaterialSerializer(self.notok_with_legacy_sign).data self.assertEqual( normal_notok_data["defect_grade"], Defect.DEFECT_NOTOK @@ -154,8 +171,14 @@ class WMaterialDefectGradeTests(TestCase): notok_b_data["defect_grade"], Defect.DEFECT_OK_B ) self.assertEqual(notok_b_data["defect_grade_name"], "合格B类") - self.assertIsNone(no_defect_data["defect_grade"]) - self.assertIsNone(no_defect_data["defect_grade_name"]) + self.assertEqual(no_defect_data["defect_grade"], Defect.DEFECT_OK) + self.assertEqual(no_defect_data["defect_grade_name"], "合格") + self.assertEqual( + repair_no_defect_data["defect_grade"], Defect.DEFECT_OK + ) + self.assertEqual(repair_no_defect_data["defect_grade_name"], "合格") + self.assertEqual(legacy_data["defect_grade"], Defect.DEFECT_NOTOK) + self.assertEqual(legacy_data["defect_grade_name"], "不合格") def test_filtering_state_and_defect_grade_are_independent(self): normal_notok = WMaterialFilter( @@ -184,6 +207,28 @@ class WMaterialDefectGradeTests(TestCase): transform=lambda item: item, ) + def test_effective_grade_filter_includes_defaults_and_legacy_signs(self): + ok_items = WMaterialFilter( + {"defect_grade": Defect.DEFECT_OK}, + queryset=WMaterial.objects.all(), + ).qs + notok_items = WMaterialFilter( + {"defect_grade": Defect.DEFECT_NOTOK}, + queryset=WMaterial.objects.all(), + ).qs + + self.assertCountEqual( + ok_items.values_list("id", flat=True), + [self.normal_without_defect.id, self.repair_without_defect.id], + ) + self.assertCountEqual( + notok_items.values_list("id", flat=True), + [ + self.normal_with_notok_defect.id, + self.notok_with_legacy_sign.id, + ], + ) + class MlogbwViewSetTests(SimpleTestCase): @patch("apps.wpm.views.MlogViewSet.lock_and_check_can_update") diff --git a/apps/wpm/tests/test_defect_grade.py b/apps/wpm/tests/test_defect_grade.py new file mode 100644 index 00000000..d2fdbf26 --- /dev/null +++ b/apps/wpm/tests/test_defect_grade.py @@ -0,0 +1,86 @@ +from django.test import TestCase + +from apps.mtm.models import Material +from apps.qm.models import Defect +from apps.wpm.filters import WMaterialFilter +from apps.wpm.models import WMaterial +from apps.wpm.serializers import WMaterialSerializer + + +class WMaterialDefectGradeTests(TestCase): + @classmethod + def setUpTestData(cls): + cls.material = Material.objects.create(name="缺陷等级测试物料") + cls.defect_b = Defect.objects.create( + name="B类缺陷", + cate=Defect.cate_list[0], + okcate=Defect.DEFECT_OK_B, + ) + cls.defect_notok = Defect.objects.create( + name="不合格缺陷", + cate=Defect.cate_list[0], + okcate=Defect.DEFECT_NOTOK, + ) + cls.repair_without_defect = WMaterial.objects.create( + material=cls.material, + batch="REPAIR-NONE", + count=1, + state=WMaterial.WM_REPAIR, + ) + cls.normal_with_notok_defect = WMaterial.objects.create( + material=cls.material, + batch="NORMAL-NOTOK", + count=1, + state=WMaterial.WM_OK, + defect=cls.defect_notok, + ) + cls.notok_with_b_defect = WMaterial.objects.create( + material=cls.material, + batch="NOTOK-B", + count=1, + state=WMaterial.WM_NOTOK, + defect=cls.defect_b, + ) + cls.notok_with_legacy_sign = WMaterial.objects.create( + material=cls.material, + batch="NOTOK-LEGACY", + count=1, + state=WMaterial.WM_NOTOK, + notok_sign="zw", + ) + + def test_serializer_uses_effective_grade_independent_of_state(self): + repair_data = WMaterialSerializer(self.repair_without_defect).data + notok_data = WMaterialSerializer( + self.normal_with_notok_defect + ).data + b_data = WMaterialSerializer(self.notok_with_b_defect).data + legacy_data = WMaterialSerializer(self.notok_with_legacy_sign).data + + self.assertEqual(repair_data["defect_grade"], Defect.DEFECT_OK) + self.assertEqual(repair_data["defect_grade_name"], "合格") + self.assertEqual(notok_data["defect_grade"], Defect.DEFECT_NOTOK) + self.assertEqual(b_data["defect_grade"], Defect.DEFECT_OK_B) + self.assertEqual(legacy_data["defect_grade"], Defect.DEFECT_NOTOK) + + def test_effective_grade_filter_matches_serializer_rules(self): + ok_items = WMaterialFilter( + {"defect_grade": Defect.DEFECT_OK}, + queryset=WMaterial.objects.all(), + ).qs + notok_items = WMaterialFilter( + {"defect_grade": Defect.DEFECT_NOTOK}, + queryset=WMaterial.objects.all(), + ).qs + + self.assertCountEqual( + ok_items.values_list("id", flat=True), + [self.repair_without_defect.id], + ) + self.assertCountEqual( + notok_items.values_list("id", flat=True), + [ + self.normal_with_notok_defect.id, + self.notok_with_legacy_sign.id, + ], + ) From fd4de2bd4b049099ec8c9d685d62e5d573052fd3 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Fri, 7 Aug 2026 09:58:54 +0800 Subject: [PATCH 07/15] fix(wpm): match historical numbers to active rule --- apps/wpm/tests/test_number_rule.py | 141 +++++++++++++++++++++++++++-- apps/wpm/views.py | 60 +++++++++--- 2 files changed, 182 insertions(+), 19 deletions(-) diff --git a/apps/wpm/tests/test_number_rule.py b/apps/wpm/tests/test_number_rule.py index ee8e7d5e..0715db3c 100644 --- a/apps/wpm/tests/test_number_rule.py +++ b/apps/wpm/tests/test_number_rule.py @@ -29,9 +29,7 @@ class GenNumberWithRuleFilterTests(SimpleTestCase): for rule, expected_dates in cases: with self.subTest(rule=rule): queryset = MagicMock() - queryset.annotate.return_value = queryset - queryset.order_by.return_value = queryset - queryset.last.return_value = None + queryset.values_list.return_value.distinct.return_value.iterator.return_value = [] with patch("apps.wpmw.models.Wpr.objects.filter", return_value=queryset) as mock_filter: MlogbInViewSet.gen_number_with_rule(rule, material, mlog) @@ -47,9 +45,7 @@ class GenNumberWithRuleFilterTests(SimpleTestCase): def test_escaped_date_placeholder_text_does_not_add_filter(self): queryset = MagicMock() - queryset.annotate.return_value = queryset - queryset.order_by.return_value = queryset - queryset.last.return_value = None + queryset.values_list.return_value.distinct.return_value.iterator.return_value = [] material = SimpleNamespace(model=None) mlog = SimpleNamespace( handle_date=date(2026, 8, 4), @@ -62,3 +58,136 @@ class GenNumberWithRuleFilterTests(SimpleTestCase): self.assertFalse( any("handle_date" in key for key in mock_filter.call_args.kwargs) ) + + def test_previous_sequence_width_is_compatible(self): + queryset = MagicMock() + queryset.values_list.return_value.distinct.return_value.iterator.return_value = [ + "2608P0001", + "2608P0003", + "2608P0002", + ] + material = SimpleNamespace(model="P") + mlog = SimpleNamespace( + handle_date=date(2026, 8, 6), + mgroup=SimpleNamespace(process=SimpleNamespace(id=123)), + ) + + with patch("apps.wpmw.models.Wpr.objects.filter", return_value=queryset): + number = MlogbInViewSet.gen_number_with_rule( + "{c_year2}{c_month:02d}{m_model}{n_count:05d}", + material, + mlog, + ) + + self.assertEqual(number, "2608P00004") + + def test_mixed_sequence_widths_use_numeric_maximum(self): + queryset = MagicMock() + queryset.values_list.return_value.distinct.return_value.iterator.return_value = [ + "2608P9999", + "2608P10000", + "历史异常编号", + ] + material = SimpleNamespace(model="P") + mlog = SimpleNamespace( + handle_date=date(2026, 8, 6), + mgroup=SimpleNamespace(process=SimpleNamespace(id=123)), + ) + + with patch("apps.wpmw.models.Wpr.objects.filter", return_value=queryset): + number = MlogbInViewSet.gen_number_with_rule( + "{c_year2}{c_month:02d}{m_model}{n_count:05d}", + material, + mlog, + ) + + self.assertEqual(number, "2608P10001") + + def test_unrelated_historical_formats_do_not_participate(self): + queryset = MagicMock() + queryset.values_list.return_value.distinct.return_value.iterator.return_value = [ + "202505508001", + "3p05013", + "05002", + "3pb003", + ] + material = SimpleNamespace(model="P") + mlog = SimpleNamespace( + handle_date=date(2026, 8, 7), + mgroup=SimpleNamespace(process=SimpleNamespace(id=123)), + ) + + with patch("apps.wpmw.models.Wpr.objects.filter", return_value=queryset): + number = MlogbInViewSet.gen_number_with_rule( + "{m_model}{n_count:04d}", + material, + mlog, + ) + + self.assertEqual(number, "P0001") + + def test_current_rule_uses_only_matching_model_numbers(self): + queryset = MagicMock() + queryset.values_list.return_value.distinct.return_value.iterator.return_value = [ + "P0002", + "P0010", + "B0099", + "3pb100", + ] + material = SimpleNamespace(model="P") + mlog = SimpleNamespace( + handle_date=date(2026, 8, 7), + mgroup=SimpleNamespace(process=SimpleNamespace(id=123)), + ) + + with patch("apps.wpmw.models.Wpr.objects.filter", return_value=queryset): + number = MlogbInViewSet.gen_number_with_rule( + "{m_model}{n_count:04d}", + material, + mlog, + ) + + self.assertEqual(number, "P0011") + + def test_empty_historical_number_is_ignored(self): + queryset = MagicMock() + queryset.values_list.return_value.distinct.return_value.iterator.return_value = [ + None, + "", + "P0002", + ] + material = SimpleNamespace(model="P") + mlog = SimpleNamespace( + handle_date=date(2026, 8, 7), + mgroup=SimpleNamespace(process=SimpleNamespace(id=123)), + ) + + with patch("apps.wpmw.models.Wpr.objects.filter", return_value=queryset): + number = MlogbInViewSet.gen_number_with_rule( + "{m_model}{n_count:04d}", + material, + mlog, + ) + + self.assertEqual(number, "P0003") + + def test_repeated_sequence_placeholder_does_not_break_matching(self): + queryset = MagicMock() + queryset.values_list.return_value.distinct.return_value.iterator.return_value = [ + "P02-0002", + "P03-0004", + ] + material = SimpleNamespace(model="P") + mlog = SimpleNamespace( + handle_date=date(2026, 8, 7), + mgroup=SimpleNamespace(process=SimpleNamespace(id=123)), + ) + + with patch("apps.wpmw.models.Wpr.objects.filter", return_value=queryset): + number = MlogbInViewSet.gen_number_with_rule( + "{m_model}{n_count:02d}-{n_count:04d}", + material, + mlog, + ) + + self.assertEqual(number, "P03-0003") diff --git a/apps/wpm/views.py b/apps/wpm/views.py index 1084145f..3c66404c 100644 --- a/apps/wpm/views.py +++ b/apps/wpm/views.py @@ -74,7 +74,6 @@ from django.db.models import Prefetch from drf_yasg.utils import swagger_auto_schema from drf_yasg import openapi from django.db import connection -from django.db.models.functions import Substr, Length from apps.qm.models import FtestDefect, FtestItem # Create your views here. @@ -1012,9 +1011,11 @@ class MlogbInViewSet(BulkCreateModelMixin, BulkUpdateModelMixin, BulkDestroyMode def gen_number_with_rule(cls, rule, material_out: Material, mlog: Mlog, gen_count=1): from apps.wpmw.models import Wpr + formatter = Formatter() + rule_parts = list(formatter.parse(rule)) rule_fields = { field_name - for _, field_name, _, _ in Formatter().parse(rule) + for _, field_name, _, _ in rule_parts if field_name } handle_date = mlog.handle_date @@ -1049,18 +1050,51 @@ class MlogbInViewSet(BulkCreateModelMixin, BulkUpdateModelMixin, BulkDestroyMode wpr_filter["wpr_mlogbw__mlogb__mlog__handle_date__month"] = c_month if "c_day" in rule_fields: wpr_filter["wpr_mlogbw__mlogb__mlog__handle_date__day"] = c_day - wpr = ( - Wpr.objects.filter(**wpr_filter) - .annotate(last_seq=Substr("number", Length("number") - (cq_w - 1))) - .order_by("last_seq") - .last() - ) - n_count = 0 - if wpr: + rule_values = { + "c_year": c_year, + "c_year2": c_year2, + "c_month": c_month, + "c_day": c_day, + "m_model": m_model, + } + number_pattern_parts = ["^"] + sequence_group_names = [] + for literal_text, field_name, format_spec, conversion in rule_parts: + number_pattern_parts.append(re.escape(literal_text)) + if not field_name: + continue + if field_name == "n_count": + # 流水号宽度可以变化,规则中的其他部分必须与当前上下文一致。 + group_name = f"n_count_{len(sequence_group_names)}" + sequence_group_names.append(group_name) + number_pattern_parts.append(fr"(?P<{group_name}>[0-9]+)") + continue try: - n_count = int(wpr.number[-cq_w:]) - except Exception as e: - raise ParseError(f"获取该类产品最后编号错误: {str(e)}") + field_value = rule_values[field_name] + if conversion: + field_value = formatter.convert_field(field_value, conversion) + formatted_value = formatter.format_field(field_value, format_spec) + except (KeyError, TypeError, ValueError) as e: + raise ParseError(f"个号生成错误: {e}") + number_pattern_parts.append(re.escape(formatted_value)) + number_pattern_parts.append("$") + number_pattern = re.compile("".join(number_pattern_parts)) + n_count = 0 + # 只从符合当前规则固定部分的历史编号中提取流水号。例如当前规则为 + # P{n_count:04d}时,3pb003等同工序的旧格式编号不能参与续号;同时 + # 流水号使用数字匹配,以兼容04d调整为05d后的历史编号。 + numbers = Wpr.objects.filter(**wpr_filter).values_list("number", flat=True).distinct() + for number in numbers.iterator(): + if not isinstance(number, str): + continue + sequence_match = number_pattern.fullmatch(number) + if sequence_match and sequence_group_names: + sequence_values = { + int(sequence_match.group(group_name)) + for group_name in sequence_group_names + } + if len(sequence_values) == 1: + n_count = max(n_count, sequence_values.pop()) if n_count + gen_count > 10 ** cq_w - 1: raise ParseError(f"流水号超出{cq_w}位上限, 请调整编号规则") try: From d98b7fada2c70d7396bc8b104f5055e45a4d561e Mon Sep 17 00:00:00 2001 From: caoqianming Date: Fri, 7 Aug 2026 10:28:12 +0800 Subject: [PATCH 08/15] fix(wpm): allow qualified and grade b batch merges --- apps/wpm/serializers.py | 12 +++------ apps/wpm/tests.py | 60 ++++++++++++++++++++++++++++++----------- 2 files changed, 47 insertions(+), 25 deletions(-) diff --git a/apps/wpm/serializers.py b/apps/wpm/serializers.py index bd868f85..2b360925 100644 --- a/apps/wpm/serializers.py +++ b/apps/wpm/serializers.py @@ -1419,7 +1419,6 @@ class HandoverSerializer(CustomModelSerializer): next_mat = None next_state = None next_defect = None - next_defect_grade = None if new_wm and attrs["type"] != Handover.H_CHANGE: next_mat = new_wm.material next_state = new_wm.state @@ -1440,15 +1439,10 @@ class HandoverSerializer(CustomModelSerializer): if clear_defect and new_wm is not None and new_wm.defect is not None: raise ParseError('清除批次缺陷时目标批次不能带缺陷') if clear_defect and tracking == Material.MA_TRACKING_BATCH: - if wm.defect is None: + defect_grade = effective_defect_grade(wm, "notok_sign") + if defect_grade not in [Defect.DEFECT_OK, Defect.DEFECT_OK_B]: raise ParseError( - f'第{ind+1}行-批次追踪物料仅同缺陷等级可清除批次缺陷' - ) - if next_defect_grade is None: - next_defect_grade = wm.defect.okcate - elif next_defect_grade != wm.defect.okcate: - raise ParseError( - f'第{ind+1}行-批次追踪物料仅同缺陷等级可清除批次缺陷' + f'第{ind+1}行-批次追踪物料仅合格品和合格B类可清除批次缺陷' ) if next_mat is None: next_mat = wm.material diff --git a/apps/wpm/tests.py b/apps/wpm/tests.py index 9ffaf96f..30d801cb 100644 --- a/apps/wpm/tests.py +++ b/apps/wpm/tests.py @@ -507,7 +507,34 @@ class WMaterialScopeTests(SimpleTestCase): self.assertTrue(validated["clear_defect"]) self.assertEqual(validated["count"], 2) - def test_batch_tracking_merge_can_clear_same_grade_notok_defects(self): + def test_batch_tracking_merge_can_clear_ok_and_ok_b_defects(self): + material = Material(tracking=Material.MA_TRACKING_BATCH) + defect_b = Defect(id="1", okcate=Defect.DEFECT_OK_B) + wm_ok = WMaterial( + id="10", material=material, batch="OK-001", count=1, + state=WMaterial.WM_OK, defect=None, + ) + wm_b = WMaterial( + id="20", material=material, batch="B-001", count=1, + state=WMaterial.WM_OK, defect=defect_b, + ) + + validated = HandoverSerializer().validate({ + "wm": wm_ok, + "handoverb": [ + {"wm": wm_ok, "count": 1}, + {"wm": wm_b, "count": 1}, + ], + "new_batch": "OK-MERGED", + "clear_defect": True, + "type": Handover.H_NORMAL, + "mtype": Handover.H_MERGE, + }) + + self.assertTrue(validated["clear_defect"]) + self.assertEqual(validated["count"], 2) + + def test_batch_tracking_merge_cannot_clear_same_grade_notok_defects(self): material = Material(tracking=Material.MA_TRACKING_BATCH) defect_a = Defect(id="1", okcate=Defect.DEFECT_NOTOK) defect_b = Defect(id="2", okcate=Defect.DEFECT_NOTOK) @@ -520,20 +547,21 @@ class WMaterialScopeTests(SimpleTestCase): state=WMaterial.WM_NOTOK, defect=defect_b, ) - validated = HandoverSerializer().validate({ - "wm": wm_a, - "handoverb": [ - {"wm": wm_a, "count": 1}, - {"wm": wm_b, "count": 1}, - ], - "new_batch": "N-MERGED", - "clear_defect": True, - "type": Handover.H_NORMAL, - "mtype": Handover.H_MERGE, - }) - - self.assertTrue(validated["clear_defect"]) - self.assertEqual(validated["count"], 2) + with self.assertRaisesMessage( + ParseError, + "批次追踪物料仅合格品和合格B类可清除批次缺陷", + ): + HandoverSerializer().validate({ + "wm": wm_a, + "handoverb": [ + {"wm": wm_a, "count": 1}, + {"wm": wm_b, "count": 1}, + ], + "new_batch": "N-MERGED", + "clear_defect": True, + "type": Handover.H_NORMAL, + "mtype": Handover.H_MERGE, + }) def test_batch_tracking_merge_cannot_clear_mixed_defect_grades(self): material = Material(tracking=Material.MA_TRACKING_BATCH) @@ -550,7 +578,7 @@ class WMaterialScopeTests(SimpleTestCase): with self.assertRaisesMessage( ParseError, - "批次追踪物料仅同缺陷等级可清除批次缺陷", + "批次追踪物料仅合格品和合格B类可清除批次缺陷", ): HandoverSerializer().validate({ "wm": wm_a, From 24ce008d3a16f57a8b277c0002c0d894474dae6d Mon Sep 17 00:00:00 2001 From: caoqianming Date: Fri, 7 Aug 2026 11:15:22 +0800 Subject: [PATCH 09/15] perf(wpm): batch load handover inventory validation --- apps/wpm/serializers.py | 40 +++++++++++++++ apps/wpm/tests/test_handover_serializer.py | 58 ++++++++++++++++++++++ 2 files changed, 98 insertions(+) create mode 100644 apps/wpm/tests/test_handover_serializer.py diff --git a/apps/wpm/serializers.py b/apps/wpm/serializers.py index 2b360925..409269aa 100644 --- a/apps/wpm/serializers.py +++ b/apps/wpm/serializers.py @@ -1265,7 +1265,26 @@ class Handoverbwserializer(CustomModelSerializer): read_only_fields = EXCLUDE_FIELDS_BASE + ["handoverb", "number"] extra_kwargs = {'wpr': {'required': True}} + +class CachedWMaterialPrimaryKeyRelatedField(serializers.PrimaryKeyRelatedField): + def to_internal_value(self, data): + cache = getattr(self.root, "_handover_wmaterial_cache", None) + if cache is None: + return super().to_internal_value(data) + if not isinstance(data, (str, int)): + self.fail("incorrect_type", data_type=type(data).__name__) + try: + return cache[str(data)] + except KeyError: + self.fail("does_not_exist", pk_value=data) + + class HandoverbSerializer(CustomModelSerializer): + wm = CachedWMaterialPrimaryKeyRelatedField( + queryset=WMaterial.objects.select_related( + "material", "defect", "mgroup", "belong_dept" + ) + ) notok_sign = serializers.CharField(source='wm.notok_sign', read_only=True) defect_name = serializers.CharField(source="wm.defect.name", read_only=True) handoverbw = Handoverbwserializer(many=True, required=False) @@ -1301,6 +1320,27 @@ class HandoverSerializer(CustomModelSerializer): wm_notok_sign = serializers.CharField(source='wm.notok_sign', read_only=True) handoverb = HandoverbSerializer(many=True, required=False) ticket_ = TicketSimpleSerializer(source='ticket', read_only=True) + + def to_internal_value(self, data): + handoverb = data.get("handoverb", []) if hasattr(data, "get") else [] + wm_ids = { + str(item["wm"]) + for item in handoverb + if isinstance(item, dict) and item.get("wm") is not None + } + if not wm_ids: + return super().to_internal_value(data) + + queryset = WMaterial.objects.select_related( + "material", "defect", "mgroup", "belong_dept" + ) + self._handover_wmaterial_cache = { + str(pk): instance for pk, instance in queryset.in_bulk(wm_ids).items() + } + try: + return super().to_internal_value(data) + finally: + del self._handover_wmaterial_cache def validate(self, attrs): if "mtype" not in attrs: diff --git a/apps/wpm/tests/test_handover_serializer.py b/apps/wpm/tests/test_handover_serializer.py new file mode 100644 index 00000000..cd73a488 --- /dev/null +++ b/apps/wpm/tests/test_handover_serializer.py @@ -0,0 +1,58 @@ +from django.db import connection +from django.test import TestCase +from django.test.utils import CaptureQueriesContext + +from apps.mtm.models import Material +from apps.qm.models import Defect +from apps.system.models import Dept, User +from apps.wpm.models import Handover, WMaterial +from apps.wpm.serializers import HandoverSerializer + + +class HandoverSerializerQueryTests(TestCase): + @classmethod + def setUpTestData(cls): + cls.dept = Dept.objects.create(name="合批查询测试车间") + cls.user = User.objects.create_user(username="handover-query-user") + cls.material = Material.objects.create(name="合批查询测试物料") + cls.defect_b = Defect.objects.create( + name="合批查询测试B类缺陷", + cate=Defect.cate_list[0], + okcate=Defect.DEFECT_OK_B, + ) + cls.inventories = [ + WMaterial.objects.create( + material=cls.material, + batch=f"QUERY-{index}", + count=1, + state=WMaterial.WM_OK, + defect=cls.defect_b if index == 2 else None, + belong_dept=cls.dept, + ) + for index in range(3) + ] + + def test_handover_inventory_is_loaded_in_one_query(self): + serializer = HandoverSerializer(data={ + "send_date": "2026-08-07", + "send_user": self.user.id, + "send_dept": self.dept.id, + "recive_dept": self.dept.id, + "handoverb": [ + {"wm": inventory.id, "count": 1} + for inventory in self.inventories + ], + "new_batch": "QUERY-MERGED", + "clear_defect": True, + "type": Handover.H_NORMAL, + "mtype": Handover.H_MERGE, + }) + + with CaptureQueriesContext(connection) as queries: + self.assertTrue(serializer.is_valid(), serializer.errors) + + inventory_queries = [ + query["sql"] for query in queries + if 'FROM "wpm_wmaterial"' in query["sql"] + ] + self.assertEqual(len(inventory_queries), 1, inventory_queries) From c17552e0ea898ae18a475f3e731b7bccca18adab Mon Sep 17 00:00:00 2001 From: caoqianming Date: Fri, 7 Aug 2026 11:15:39 +0800 Subject: [PATCH 10/15] test: isolate database settings from production --- .codex/memory/MEMORY.md | 1 + .codex/memory/reference_test_database.md | 15 +++++++++++++++ manage.py | 7 ++++++- pytest.ini | 3 +++ server/test_settings.py | 16 ++++++++++++++++ 5 files changed, 41 insertions(+), 1 deletion(-) create mode 100644 .codex/memory/reference_test_database.md create mode 100644 pytest.ini create mode 100644 server/test_settings.py diff --git a/.codex/memory/MEMORY.md b/.codex/memory/MEMORY.md index 6bf064fe..c68b79e3 100644 --- a/.codex/memory/MEMORY.md +++ b/.codex/memory/MEMORY.md @@ -8,6 +8,7 @@ - [合批原料字段历史问题](project_material_ofrom_merge_bug.md):`material_ofrom` 不一致的既有排查结论。 - [前端独立发版流程](reference_ehs_web_release.md):发布 `ehs_web` 时使用。 - [项目 Python 虚拟环境](reference_python_venv.md):运行 Django、pytest 和脚本时使用。 +- [项目测试数据库](reference_test_database.md):测试与可切换的生产查询连接解耦,始终使用固定的 `test_ehs_develop`。 - [两个 WebView 套壳 App](reference_wrapper_apps.md):修改 h5x 与扫码、返回键交互时使用。 这些文件记录的是长期约定或历史上下文。执行任务前应结合当前代码和数据重新验证,尤其不要把历史缺陷结论直接当成当前故障原因。 diff --git a/.codex/memory/reference_test_database.md b/.codex/memory/reference_test_database.md new file mode 100644 index 00000000..3bb5085c --- /dev/null +++ b/.codex/memory/reference_test_database.md @@ -0,0 +1,15 @@ +# 项目测试数据库 + +- 生产问题查询时,默认数据库连接可以根据工厂或环境切换到不同 IP 和业务库;这类连接只用于授权范围内的生产数据只读验证。 +- 所有 Django、pytest 及其他自动化测试必须与当前生产查询连接解耦,始终使用固定的 `test_ehs_develop` 测试数据库连接。 +- 不能仅依赖当前默认连接的 Django 自动 `test_` 命名;即使当前业务库是 `bxerp` 或其他库,测试也不应转而使用 `test_bxerp` 或其他派生库。 +- 不得在生产数据库上运行会建表、迁移、写入或清理数据的测试。 +- 生产查询与固定测试库的连接参数均从本机已忽略配置中读取,不在项目记忆、源码或提交信息中记录凭据。 + +当前实现: + +- 固定测试连接保存在已忽略的 `config/conf_test.py`。 +- `server/test_settings.py` 加载该连接,`manage.py test` 会自动选择测试 settings;其他 `manage.py` 命令仍使用当前业务库连接。 +- 测试 settings 中的基础 `NAME` 和 `TEST.NAME` 都必须是 `test_ehs_develop`,并在启动时校验,防止通过测试 settings 误操作 `ehs_develop` 或任何生产库。 +- pytest-django 通过项目根目录 `pytest.ini` 固定使用 `server.test_settings`。 +- 本地重复运行测试时优先使用 `.venv\\Scripts\\python.exe manage.py test --keepdb --noinput`。 diff --git a/manage.py b/manage.py index 1c818788..039a3cf0 100755 --- a/manage.py +++ b/manage.py @@ -5,7 +5,12 @@ import sys def main(): - os.environ.setdefault('DJANGO_SETTINGS_MODULE', 'server.settings') + settings_module = ( + 'server.test_settings' + if sys.argv[1:2] == ['test'] + else 'server.settings' + ) + os.environ.setdefault('DJANGO_SETTINGS_MODULE', settings_module) try: from django.core.management import execute_from_command_line except ImportError as exc: diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 00000000..daee44fb --- /dev/null +++ b/pytest.ini @@ -0,0 +1,3 @@ +[pytest] +DJANGO_SETTINGS_MODULE = server.test_settings +python_files = tests.py test_*.py *_tests.py diff --git a/server/test_settings.py b/server/test_settings.py new file mode 100644 index 00000000..9fec0e2a --- /dev/null +++ b/server/test_settings.py @@ -0,0 +1,16 @@ +from server.settings import * # noqa: F403 + +from config.conf_test import TEST_DATABASES +from django.core.exceptions import ImproperlyConfigured + + +DATABASES = TEST_DATABASES + +if ( + DATABASES["default"].get("NAME") != "test_ehs_develop" + or DATABASES["default"].get("TEST", {}).get("NAME") + != "test_ehs_develop" +): + raise ImproperlyConfigured( + "Test settings must only use the test_ehs_develop database." + ) From 94215df0b69ac30abb8a7ade94773cd735bff331 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Fri, 7 Aug 2026 14:29:56 +0800 Subject: [PATCH 11/15] feat(wpm): enrich sub-process operation records --- apps/em/filters.py | 2 ++ apps/wpm/migrations/0139_mloguser_note.py | 16 ++++++++++++++++ apps/wpm/models.py | 1 + 3 files changed, 19 insertions(+) create mode 100644 apps/wpm/migrations/0139_mloguser_note.py diff --git a/apps/em/filters.py b/apps/em/filters.py index 443e6f74..fb470786 100644 --- a/apps/em/filters.py +++ b/apps/em/filters.py @@ -6,6 +6,8 @@ from apps.utils.filters import MyJsonListFilter class EquipFilterSet(filters.FilterSet): tags = MyJsonListFilter(label='tags/json/list查询') + exclude_cate_name = filters.CharFilter( + field_name='cate__name', exclude=True, label='排除设备分类名称') class Meta: model = Equipment diff --git a/apps/wpm/migrations/0139_mloguser_note.py b/apps/wpm/migrations/0139_mloguser_note.py new file mode 100644 index 00000000..49059e92 --- /dev/null +++ b/apps/wpm/migrations/0139_mloguser_note.py @@ -0,0 +1,16 @@ +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ('wpm', '0138_mlogbw_files'), + ] + + operations = [ + migrations.AddField( + model_name='mloguser', + name='note', + field=models.TextField(blank=True, default='', verbose_name='备注'), + ), + ] diff --git a/apps/wpm/models.py b/apps/wpm/models.py index 9a3c0e51..f4c87c76 100644 --- a/apps/wpm/models.py +++ b/apps/wpm/models.py @@ -550,6 +550,7 @@ class MlogUser(BaseModel): Equipment, verbose_name='生产设备', on_delete=models.CASCADE, null=True, blank=True, related_name='mloguser_equipment') shift = models.ForeignKey(Shift, verbose_name='关联班次', on_delete=models.CASCADE) handle_date = models.DateField('操作日期') + note = models.TextField('备注', default='', blank=True) class Mlogb(BaseModel): """ From fa16694f9f56f6832385bb71c759821439273b39 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Fri, 7 Aug 2026 16:04:17 +0800 Subject: [PATCH 12/15] feat(wpm): enrich handover detail records --- apps/wpm/models.py | 2 +- apps/wpm/serializers.py | 47 +++++++++++++++++++++- apps/wpm/tests/test_handover_serializer.py | 31 +++++++++++++- apps/wpm/views.py | 17 +++++++- 4 files changed, 92 insertions(+), 5 deletions(-) diff --git a/apps/wpm/models.py b/apps/wpm/models.py index f4c87c76..e2ae0105 100644 --- a/apps/wpm/models.py +++ b/apps/wpm/models.py @@ -877,7 +877,7 @@ class Handoverb(BaseModel): @property def handoverbw(self): - return Handoverbw.objects.filter(handoverb=self) + return self.w_handoverb.all() class Handoverbw(BaseModel): """TN: 单个产品交接记录 diff --git a/apps/wpm/serializers.py b/apps/wpm/serializers.py index 409269aa..91bfe6de 100644 --- a/apps/wpm/serializers.py +++ b/apps/wpm/serializers.py @@ -31,6 +31,16 @@ from apps.em.models import Equipment from django.db.models import Q mylogger = logging.getLogger("log") +WM_STATE_NAMES = { + WMaterial.WM_OK: "合格", + WMaterial.WM_NOTOK: "不合格", + WMaterial.WM_REPAIR: "返修", + WMaterial.WM_REPAIRED: "返修完成", + WMaterial.WM_TEST: "检验", + WMaterial.WM_SCRAP: "报废", +} + + class OtherLogSerializer(CustomModelSerializer): class Meta: model = OtherLog @@ -1286,8 +1296,35 @@ class HandoverbSerializer(CustomModelSerializer): ) ) notok_sign = serializers.CharField(source='wm.notok_sign', read_only=True) + notok_sign_name = serializers.SerializerMethodField() defect_name = serializers.CharField(source="wm.defect.name", read_only=True) + defect_grade = serializers.SerializerMethodField() + defect_grade_name = serializers.SerializerMethodField() + material_name = serializers.StringRelatedField(source="wm.material", read_only=True) + state_name = serializers.SerializerMethodField() + mgroup_name = serializers.CharField(source="wm.mgroup.name", read_only=True) + belong_dept_name = serializers.CharField(source="wm.belong_dept.name", read_only=True) + count_available = serializers.SerializerMethodField() handoverbw = Handoverbwserializer(many=True, required=False) + + def get_notok_sign_name(self, obj): + return getattr(NotOkOption, obj.wm.notok_sign, NotOkOption.qt).label if obj.wm.notok_sign else None + + def get_defect_grade(self, obj): + return effective_defect_grade(obj.wm, "notok_sign") + + def get_defect_grade_name(self, obj): + return DEFECT_GRADE_NAMES[self.get_defect_grade(obj)] + + def get_state_name(self, obj): + return WM_STATE_NAMES.get(obj.wm.state, str(obj.wm.state)) + + def get_count_available(self, obj): + # 编辑未提交交接时,当前明细占用的数量仍应允许重新填写。 + if obj.handover.submit_time is not None: + return obj.count + return obj.wm.count - obj.wm.count_handovering + obj.count + class Meta: model = Handoverb fields = "__all__" @@ -1310,9 +1347,14 @@ class HandoverSerializer(CustomModelSerializer): recive_user_name = serializers.CharField( source='recive_user.name', read_only=True) recive_dept_name = serializers.CharField( - source='recive_dept', read_only=True) + source='recive_dept.name', read_only=True) + send_dept_name = serializers.CharField(source='send_dept.name', read_only=True) send_mgroup_name = serializers.CharField(source='send_mgroup.name', read_only=True) recive_mgroup_name = serializers.CharField(source='recive_mgroup.name', read_only=True) + submit_user_name = serializers.CharField(source='submit_user.name', read_only=True) + type_name = serializers.CharField(source='get_type_display', read_only=True) + mtype_name = serializers.CharField(source='get_mtype_display', read_only=True) + state_changed_name = serializers.SerializerMethodField() material_ = MaterialSimpleSerializer(source='material', read_only=True) material_name = serializers.StringRelatedField( source='material', read_only=True) @@ -1321,6 +1363,9 @@ class HandoverSerializer(CustomModelSerializer): handoverb = HandoverbSerializer(many=True, required=False) ticket_ = TicketSimpleSerializer(source='ticket', read_only=True) + def get_state_changed_name(self, obj): + return WM_STATE_NAMES.get(obj.state_changed) if obj.state_changed is not None else None + def to_internal_value(self, data): handoverb = data.get("handoverb", []) if hasattr(data, "get") else [] wm_ids = { diff --git a/apps/wpm/tests/test_handover_serializer.py b/apps/wpm/tests/test_handover_serializer.py index cd73a488..40d930e7 100644 --- a/apps/wpm/tests/test_handover_serializer.py +++ b/apps/wpm/tests/test_handover_serializer.py @@ -5,7 +5,7 @@ from django.test.utils import CaptureQueriesContext from apps.mtm.models import Material from apps.qm.models import Defect from apps.system.models import Dept, User -from apps.wpm.models import Handover, WMaterial +from apps.wpm.models import Handover, Handoverb, WMaterial from apps.wpm.serializers import HandoverSerializer @@ -56,3 +56,32 @@ class HandoverSerializerQueryTests(TestCase): if 'FROM "wpm_wmaterial"' in query["sql"] ] self.assertEqual(len(inventory_queries), 1, inventory_queries) + + def test_detail_contains_display_fields_for_handover_and_items(self): + handover = Handover.objects.create( + send_date="2026-08-07", + send_user=self.user, + send_dept=self.dept, + recive_dept=self.dept, + material=self.material, + wm=self.inventories[0], + count=1, + type=Handover.H_NORMAL, + mtype=Handover.H_NORMAL, + ) + Handoverb.objects.create( + handover=handover, + wm=self.inventories[0], + batch=self.inventories[0].batch, + count=1, + ) + + data = HandoverSerializer(handover).data + + self.assertEqual(data["send_dept_name"], self.dept.name) + self.assertEqual(data["recive_dept_name"], self.dept.name) + self.assertEqual(data["type_name"], "正常交接") + self.assertEqual(data["mtype_name"], "正常") + self.assertEqual(data["handoverb"][0]["material_name"], str(self.material)) + self.assertEqual(data["handoverb"][0]["state_name"], "合格") + self.assertEqual(data["handoverb"][0]["count_available"], 1) diff --git a/apps/wpm/views.py b/apps/wpm/views.py index 3c66404c..fb31d11b 100644 --- a/apps/wpm/views.py +++ b/apps/wpm/views.py @@ -592,7 +592,20 @@ class HandoverViewSet(CustomModelViewSet): select_related_fields = ["send_user", "send_mgroup", "send_dept", "recive_user", "recive_mgroup", "recive_dept", "wm", "material_changed", "material", "material__process"] filterset_class = HandoverFilter search_fields = ["material__name", "material__number", "material__specification", "batch", "material__model", "b_handover__batch", "new_batch", "wm__batch"] - prefetch_related_fields = [Prefetch("b_handover", queryset=Handoverb.objects.select_related("wm__defect")), "ticket__state"] + prefetch_related_fields = ["ticket__state"] + + def get_queryset_custom(self, queryset): + if self.action not in ["list", "retrieve"]: + return queryset + + detail_queryset = Handoverb.objects.select_related( + "handover", "wm__defect", "wm__material", "wm__mgroup", "wm__belong_dept" + ) + if self.action == "retrieve": + detail_queryset = detail_queryset.prefetch_related("w_handoverb") + return queryset.prefetch_related( + Prefetch("b_handover", queryset=detail_queryset) + ) def perform_destroy(self, instance: Handover): user = self.request.user @@ -680,7 +693,7 @@ class HandoverViewSet(CustomModelViewSet): m_qs = m_qs.filter(process__route_p__material_in__id=materialInId) | m_qs.filter(process__route_p__routemat_route__material__id=materialInId) elif type in [Handover.H_SCRAP]: m_qs = m_qs.filter(process=None) - return Response(list(m_qs.values("id", "name").distinct())) + return Response(list(m_qs.values("id", "name", "belong_dept").distinct())) @action(methods=["post"], detail=False, perms_map={"post": "handover.create"}, serializer_class=GenHandoverWmSerializer) @transaction.atomic From 596f187d3d45eb304729a97410c96bd0dc95ef96 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Fri, 7 Aug 2026 16:17:23 +0800 Subject: [PATCH 13/15] docs: record frontend validation timing --- .codex/memory/MEMORY.md | 1 + .codex/memory/feedback_frontend_validation.md | 5 +++++ 2 files changed, 6 insertions(+) create mode 100644 .codex/memory/feedback_frontend_validation.md diff --git a/.codex/memory/MEMORY.md b/.codex/memory/MEMORY.md index c68b79e3..191d6f81 100644 --- a/.codex/memory/MEMORY.md +++ b/.codex/memory/MEMORY.md @@ -10,5 +10,6 @@ - [项目 Python 虚拟环境](reference_python_venv.md):运行 Django、pytest 和脚本时使用。 - [项目测试数据库](reference_test_database.md):测试与可切换的生产查询连接解耦,始终使用固定的 `test_ehs_develop`。 - [两个 WebView 套壳 App](reference_wrapper_apps.md):修改 h5x 与扫码、返回键交互时使用。 +- [前端验证时机](feedback_frontend_validation.md):日常修改先跑 check,完整 build 留到 push 前执行。 这些文件记录的是长期约定或历史上下文。执行任务前应结合当前代码和数据重新验证,尤其不要把历史缺陷结论直接当成当前故障原因。 diff --git a/.codex/memory/feedback_frontend_validation.md b/.codex/memory/feedback_frontend_validation.md new file mode 100644 index 00000000..4837934c --- /dev/null +++ b/.codex/memory/feedback_frontend_validation.md @@ -0,0 +1,5 @@ +# 前端验证时机 + +- 修改配套前端 `../ehs_web` 时,日常开发和中间验证优先运行项目已有的 `check`,不要每次修改后都运行完整 `build`。 +- 准备 push 前运行一次完整 `build`,用于发现生产构建阶段的问题。 +- 若当前前端尚未配置 `check` 脚本,应先说明现状,不得把其他命令擅自当作 `check`。 From d1693799b9832cc20692e035c2e294d43f474b4a Mon Sep 17 00:00:00 2001 From: caoqianming Date: Fri, 7 Aug 2026 16:24:45 +0800 Subject: [PATCH 14/15] feat(wpm): recommend equipment for production logs --- apps/wpm/serializers.py | 12 +++ apps/wpm/services.py | 39 +++++++ apps/wpm/tests/test_equipment_log_options.py | 52 ++++++++++ apps/wpm/views.py | 102 ++++++++++++++++++- 4 files changed, 202 insertions(+), 3 deletions(-) create mode 100644 apps/wpm/tests/test_equipment_log_options.py diff --git a/apps/wpm/serializers.py b/apps/wpm/serializers.py index 91bfe6de..24c63327 100644 --- a/apps/wpm/serializers.py +++ b/apps/wpm/serializers.py @@ -41,6 +41,18 @@ WM_STATE_NAMES = { } +class MlogEquipmentOptionSerializer(serializers.ModelSerializer): + mgroup_name = serializers.CharField(source="mgroup.name", read_only=True) + full_name = serializers.SerializerMethodField() + + def get_full_name(self, obj): + return f"{obj.number}|{obj.name}|{obj.model}" + + class Meta: + model = Equipment + fields = ["id", "name", "number", "model", "mgroup_name", "full_name"] + + class OtherLogSerializer(CustomModelSerializer): class Meta: model = OtherLog diff --git a/apps/wpm/services.py b/apps/wpm/services.py index b86b55fe..a41addd2 100644 --- a/apps/wpm/services.py +++ b/apps/wpm/services.py @@ -1,4 +1,5 @@ import datetime +from collections import defaultdict from django.core.cache import cache from django.db.models import Sum @@ -27,6 +28,44 @@ from django.db.models import F myLogger = logging.getLogger('log') +RECENT_EQUIPMENT_LOG_LIMIT = 50 + + +def get_recent_mgroup_equipment_ids( + mgroup_id, log_limit=RECENT_EQUIPMENT_LOG_LIMIT +): + """按日志时间倒序返回工段最近使用过的设备 ID,空值和重复值忽略。""" + recent_logs = list( + Mlog.objects.filter(mgroup_id=mgroup_id) + .order_by("-create_time", "-id") + .values_list("id", "equipment_id", "equipment_2_id")[:log_limit] + ) + if not recent_logs: + return [] + + log_ids = [log_id for log_id, _, _ in recent_logs] + multiple_equipment_ids = defaultdict(list) + for log_id, equipment_id in ( + Mlog.equipments.through.objects.filter(mlog_id__in=log_ids) + .order_by("id") + .values_list("mlog_id", "equipment_id") + ): + multiple_equipment_ids[log_id].append(equipment_id) + + result = [] + seen = set() + for log_id, equipment_id, equipment_2_id in recent_logs: + candidate_ids = [ + equipment_id, + equipment_2_id, + *multiple_equipment_ids[log_id], + ] + for candidate_id in candidate_ids: + if candidate_id and candidate_id not in seen: + seen.add(candidate_id) + result.append(candidate_id) + return result + def inherit_zt_batch(source: BatchSt, target: BatchSt): """拆批/报工改号时目标批继承来源批的直通统计大批归属(纯继承, 不做判定) diff --git a/apps/wpm/tests/test_equipment_log_options.py b/apps/wpm/tests/test_equipment_log_options.py new file mode 100644 index 00000000..a5d4cca2 --- /dev/null +++ b/apps/wpm/tests/test_equipment_log_options.py @@ -0,0 +1,52 @@ +from unittest.mock import MagicMock, patch + +from django.test import SimpleTestCase + +from apps.wpm.services import get_recent_mgroup_equipment_ids + + +class RecentMgroupEquipmentTests(SimpleTestCase): + @patch("apps.wpm.services.Mlog") + def test_collects_all_equipment_fields_in_log_order_without_duplicates( + self, mlog + ): + log_queryset = MagicMock() + mlog.objects.filter.return_value = log_queryset + log_queryset.order_by.return_value.values_list.return_value.__getitem__.return_value = [ + ("log-new", "equipment-a", None), + ("log-middle", "equipment-b", "equipment-a"), + ("log-empty", None, None), + ] + through_queryset = MagicMock() + mlog.equipments.through.objects.filter.return_value = through_queryset + through_queryset.order_by.return_value.values_list.return_value = [ + ("log-new", "equipment-c"), + ("log-middle", "equipment-c"), + ("log-middle", "equipment-d"), + ] + + result = get_recent_mgroup_equipment_ids("mgroup-1", log_limit=50) + + self.assertEqual( + result, + ["equipment-a", "equipment-c", "equipment-b", "equipment-d"], + ) + log_queryset.order_by.return_value.values_list.return_value.__getitem__.assert_called_once_with( + slice(None, 50, None) + ) + + @patch("apps.wpm.services.Mlog") + def test_returns_empty_when_recent_logs_have_no_equipment(self, mlog): + log_queryset = MagicMock() + mlog.objects.filter.return_value = log_queryset + log_queryset.order_by.return_value.values_list.return_value.__getitem__.return_value = [ + ("log-1", None, None), + ("log-2", None, None), + ] + through_queryset = MagicMock() + mlog.equipments.through.objects.filter.return_value = through_queryset + through_queryset.order_by.return_value.values_list.return_value = [] + + result = get_recent_mgroup_equipment_ids("mgroup-1") + + self.assertEqual(result, []) diff --git a/apps/wpm/views.py b/apps/wpm/views.py index fb31d11b..0b2db844 100644 --- a/apps/wpm/views.py +++ b/apps/wpm/views.py @@ -7,7 +7,7 @@ from rest_framework.decorators import action from rest_framework.exceptions import ParseError from rest_framework.response import Response from rest_framework.serializers import Serializer -from django.db.models import Sum +from django.db.models import Case, IntegerField, Sum, When from django.utils import timezone from apps.system.models import User @@ -54,13 +54,23 @@ from .serializers import ( MlogUserSerializer, BatchLogSerializer, MlogQuickSerializer, + MlogEquipmentOptionSerializer, MlogbwStartTestSerializer, HandoverListSerializer, BatchChangeSerializer, MlogbOutPatchUpdateSerializer ) -from .services import mlog_submit, handover_submit, mlog_revert, get_batch_dag, handover_revert -from apps.wpm.services import mlog_submit_validate, generate_new_batch +from .services import ( + RECENT_EQUIPMENT_LOG_LIMIT, + generate_new_batch, + get_batch_dag, + get_recent_mgroup_equipment_ids, + handover_revert, + handover_submit, + mlog_revert, + mlog_submit, + mlog_submit_validate, +) from apps.wf.models import State, Ticket from apps.wpmw.models import Wpr from apps.qm.models import Qct, Ftest, TestItem @@ -332,6 +342,92 @@ class MlogViewSet(CustomModelViewSet): ] ordering_fields = ["create_time", "update_time"] + @swagger_auto_schema( + manual_parameters=[ + openapi.Parameter( + name="mgroup", + in_=openapi.IN_QUERY, + description="日志所属工段", + type=openapi.TYPE_STRING, + required=True, + ), + openapi.Parameter( + name="search", + in_=openapi.IN_QUERY, + description="按设备名称或编号搜索全部生产设备", + type=openapi.TYPE_STRING, + required=False, + ), + ] + ) + @action( + methods=["get"], + detail=False, + perms_map={"get": "*"}, + serializer_class=MlogEquipmentOptionSerializer, + ) + def equipment_options(self, request, *args, **kwargs): + """返回本工段设备、最近 50 条日志用过的设备或搜索结果。""" + mgroup_id = request.query_params.get("mgroup") + if not mgroup_id: + raise ParseError("请传入mgroup参数") + + search = request.query_params.get("search", "").strip() + owned_ids = list( + Equipment.objects.filter( + type=Equipment.EQUIP_TYPE_PRO, + mgroup_id=mgroup_id, + ) + .order_by("name", "number") + .values_list("id", flat=True) + ) + owned_id_set = set(owned_ids) + + if search: + queryset = ( + Equipment.objects.filter(type=Equipment.EQUIP_TYPE_PRO) + .filter(Q(name__icontains=search) | Q(number__icontains=search)) + .order_by("name", "number") + ) + option_group = "搜索结果" + else: + recent_ids = get_recent_mgroup_equipment_ids( + mgroup_id, RECENT_EQUIPMENT_LOG_LIMIT + ) + option_ids = list(dict.fromkeys([*owned_ids, *recent_ids])) + if option_ids: + order = Case( + *[ + When(id=equipment_id, then=position) + for position, equipment_id in enumerate(option_ids) + ], + output_field=IntegerField(), + ) + queryset = Equipment.objects.filter( + id__in=option_ids, + type=Equipment.EQUIP_TYPE_PRO, + ).order_by(order) + else: + queryset = Equipment.objects.none() + option_group = None + + queryset = queryset.select_related("mgroup") + page = self.paginate_queryset(queryset) + equipment_list = page if page is not None else queryset + data = MlogEquipmentOptionSerializer( + equipment_list, + many=True, + context=self.get_serializer_context(), + ).data + for item in data: + item["option_group"] = option_group or ( + "本工段设备" if item["id"] in owned_id_set else "近期使用" + ) + + if page is not None: + return self.get_paginated_response(data) + return Response(data) + def add_info_for_item(self, data): if data.get("oinfo_json", {}): czx_dict = dict(TestItem.objects.filter(id__in=data.get("oinfo_json", {}).keys()).values_list("id", "name")) From a7a9a0aa6f85f95f3c1751a2f6d4e7661c958103 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Mon, 10 Aug 2026 13:47:56 +0800 Subject: [PATCH 15/15] feat(mcp): add factory domain tools --- apps/bi/services.py | 86 +++++++++++------ apps/bi/test_services.py | 105 ++++++++++++++++++++ apps/bi/views.py | 61 +++--------- docs/mcp.md | 49 ++++++++++ mcp_server/__init__.py | 1 + mcp_server/__main__.py | 24 +++++ mcp_server/auth.py | 66 +++++++++++++ mcp_server/context.py | 32 +++++++ mcp_server/server.py | 89 +++++++++++++++++ mcp_server/test_batch_stats.py | 105 ++++++++++++++++++++ mcp_server/test_datasets.py | 96 +++++++++++++++++++ mcp_server/test_wprs.py | 103 ++++++++++++++++++++ mcp_server/tests.py | 164 ++++++++++++++++++++++++++++++++ mcp_server/tools/__init__.py | 1 + mcp_server/tools/batch_stats.py | 104 ++++++++++++++++++++ mcp_server/tools/common.py | 18 ++++ mcp_server/tools/datasets.py | 79 +++++++++++++++ mcp_server/tools/wprs.py | 139 +++++++++++++++++++++++++++ requirements.txt | 5 + server/settings.py | 14 +++ 20 files changed, 1263 insertions(+), 78 deletions(-) create mode 100644 apps/bi/test_services.py create mode 100644 docs/mcp.md create mode 100644 mcp_server/__init__.py create mode 100644 mcp_server/__main__.py create mode 100644 mcp_server/auth.py create mode 100644 mcp_server/context.py create mode 100644 mcp_server/server.py create mode 100644 mcp_server/test_batch_stats.py create mode 100644 mcp_server/test_datasets.py create mode 100644 mcp_server/test_wprs.py create mode 100644 mcp_server/tests.py create mode 100644 mcp_server/tools/__init__.py create mode 100644 mcp_server/tools/batch_stats.py create mode 100644 mcp_server/tools/common.py create mode 100644 mcp_server/tools/datasets.py create mode 100644 mcp_server/tools/wprs.py diff --git a/apps/bi/services.py b/apps/bi/services.py index ee896bd8..0a82c49d 100644 --- a/apps/bi/services.py +++ b/apps/bi/services.py @@ -1,11 +1,15 @@ -from rest_framework.exceptions import ParseError +import concurrent.futures import json -from jinja2 import Template +import logging + +from rest_framework.exceptions import ParseError + from apps.bi.models import Dataset -import concurrent from apps.utils.sql import execute_raw_sql, format_sqldata from apps.utils.tools import MyJSONEncoder +myLogger = logging.getLogger('log') + forbidden_keywords = ["UPDATE", "DELETE", "DROP", "TRUNCATE", "INSERT", "CREATE", "ALTER", "GRANT", "REVOKE", "EXEC", "EXECUTE"] @@ -32,31 +36,57 @@ def format_json_with_placeholders(json_str, **kwargs): return formatted_json -def exec_dataset(dt: Dataset, xquery: dict = {}): +def render_dataset_sql(dt: Dataset, xquery=None, *, is_test=False): + """根据数据集配置和调用参数生成经过安全检查的只读 SQL。""" + query = dict(dt.default_param or {}) + query.update(dict(dt.test_param or {}) if is_test else dict(xquery or {})) + if not dt.sql_query: + return '' + try: + return check_sql_safe(dt.sql_query.format(**query)) + except KeyError as exc: + raise ParseError(f'需指定查询参数_{str(exc)}') from exc + + +def execute_rendered_dataset(dt: Dataset, full_sql: str, *, raise_exception=True): + """执行已经渲染和校验的 SQL,返回可合并到数据集响应的结果。""" + results = {} + results2 = {} + can_cache = True + sql_list = [sql for sql in full_sql.strip(';').split(';') if sql.strip()] + if sql_list: + with concurrent.futures.ThreadPoolExecutor(max_workers=6) as executor: + futures = { + executor.submit(execute_raw_sql, sql): (f'ds{index}', sql) + for index, sql in enumerate(sql_list) + } + for future in concurrent.futures.as_completed(futures): + name, sql = futures[future] + try: + res = future.result() + results[name], results2[name] = format_sqldata(res[0], res[1]) + except Exception as exc: + myLogger.error(f'bi查询异常:{str(exc)}-{dt.code}--{sql}') + if raise_exception: + raise ParseError(f'查询异常:{str(exc)}') from exc + results[name] = 'error: ' + str(exc) + can_cache = False + + response_data = {'data': results, 'data2': results2} + if dt.echart_options and not dt.echart_options.startswith('function'): + for result in results.values(): + if isinstance(result, str): + raise ParseError(result) + response_data['echart_options'] = format_json_with_placeholders( + dt.echart_options, **results + ) + return response_data, can_cache + + +def exec_dataset(dt: Dataset, xquery=None): """执行数据集 返回 (sql语句, { rda}) """ - rdata = {} - results = {} - results2 = {} - query = dt.default_param - if dt.sql_query: - query.update(xquery) - sql_f_ = check_sql_safe(dt.sql_query.format(**query)) - sql_f_strip = sql_f_.strip(';') - sql_f_l = sql_f_strip.split(';') - # 多线程运行并返回字典结果 - with concurrent.futures.ThreadPoolExecutor(max_workers=6) as executor: - fun_ps = [] - for ind, val in enumerate(sql_f_l): - fun_ps.append((f'ds{ind}', execute_raw_sql, val)) - # 生成执行函数 - futures = {executor.submit(i[1], i[2]): i for i in fun_ps} - for future in concurrent.futures.as_completed(futures): - name, *_, sql_f = futures[future] # 获取对应的键 - res = future.result() - results[name], results2[name] = format_sqldata( - res[0], res[1]) - rdata['data'] = results - rdata['data2'] = results2 - return sql_f_, rdata \ No newline at end of file + full_sql = render_dataset_sql(dt, xquery) + response_data, _ = execute_rendered_dataset(dt, full_sql) + return full_sql, response_data diff --git a/apps/bi/test_services.py b/apps/bi/test_services.py new file mode 100644 index 00000000..2e18a409 --- /dev/null +++ b/apps/bi/test_services.py @@ -0,0 +1,105 @@ +from types import SimpleNamespace +from unittest.mock import patch + +from django.test import SimpleTestCase +from rest_framework.exceptions import ParseError + +from apps.bi.services import ( + exec_dataset, + execute_rendered_dataset, + render_dataset_sql, +) +from apps.bi.views import DatasetViewSet + + +def dataset(**overrides): + values = { + "code": "output_daily", + "sql_query": "select * from output where day = '{day}'", + "default_param": {"day": "2026-08-01"}, + "test_param": {"day": "2026-08-02"}, + "echart_options": "", + } + values.update(overrides) + return SimpleNamespace(**values) + + +class DatasetExecutionServiceTests(SimpleTestCase): + def test_render_does_not_mutate_default_parameters(self): + item = dataset() + + sql = render_dataset_sql(item, {"day": "2026-08-10"}) + + self.assertIn("2026-08-10", sql) + self.assertEqual(item.default_param, {"day": "2026-08-01"}) + + def test_render_reports_missing_parameters(self): + item = dataset(default_param={}, sql_query="select '{required}'") + + with self.assertRaises(ParseError): + render_dataset_sql(item) + + def test_execute_formats_each_statement(self): + item = dataset(echart_options='{"series": {ds0}}') + with ( + patch("apps.bi.services.execute_raw_sql", return_value=([], [])), + patch( + "apps.bi.services.format_sqldata", + return_value=([{"count": 1}], {"count": [1]}), + ), + ): + response, can_cache = execute_rendered_dataset( + item, "select 1;select 2" + ) + + self.assertTrue(can_cache) + self.assertEqual(set(response["data"]), {"ds0", "ds1"}) + self.assertIn('"count": 1', response["echart_options"]) + + def test_empty_dataset_has_stable_empty_result(self): + full_sql, response = exec_dataset(dataset(sql_query="")) + + self.assertEqual(full_sql, "") + self.assertEqual(response, {"data": {}, "data2": {}}) + + +class DatasetViewExecutionTests(SimpleTestCase): + @patch("apps.bi.views.cache") + @patch("apps.bi.views.execute_rendered_dataset") + @patch("apps.bi.views.render_dataset_sql") + @patch("apps.bi.views.DatasetSerializer") + def test_api_reuses_shared_execution_service( + self, + serializer_mock, + render_mock, + execute_mock, + cache_mock, + ): + item = dataset( + enabled=True, + name="日产量", + cache_seconds=10, + ) + serializer_mock.return_value.data = { + "code": item.code, + "echart_options": "", + } + render_mock.return_value = "select 1" + execute_mock.return_value = ( + {"data": {"ds0": [{"count": 1}]}, "data2": {}}, + True, + ) + cache_mock.get.return_value = None + view = DatasetViewSet(basename="dataset") + view.get_object = lambda: item + request = SimpleNamespace( + data={"query": {"day": "2026-08-10"}}, + user=SimpleNamespace(id=42, belong_dept=SimpleNamespace(id=7)), + ) + + response = view.exec(request) + + render_query = render_mock.call_args.args[1] + self.assertEqual(render_query["r_user"], 42) + self.assertEqual(render_query["r_dept"], 7) + self.assertEqual(response.data["data"]["ds0"][0]["count"], 1) diff --git a/apps/bi/views.py b/apps/bi/views.py index 7c5ac1d3..6004ce1e 100644 --- a/apps/bi/views.py +++ b/apps/bi/views.py @@ -11,17 +11,13 @@ from apps.bi.serializers import ( DatasetSerializer, ) from django.apps import apps -import concurrent.futures from django.core.cache import cache -from apps.utils.sql import execute_raw_sql, format_sqldata -from apps.bi.services import check_sql_safe, format_json_with_placeholders +from apps.bi.services import execute_rendered_dataset, render_dataset_sql from rest_framework.exceptions import ParseError from rest_framework.generics import get_object_or_404 from apps.utils.mixins import ListModelMixin -import logging from drf_yasg import openapi from drf_yasg.utils import swagger_auto_schema -myLogger = logging.getLogger('log') # Create your views here. @@ -132,59 +128,24 @@ class DatasetViewSet(CustomModelViewSet): if not dt.enabled: raise ParseError(f'{dt.name}-该查询未启用') rdata = DatasetSerializer(instance=dt).data - xquery = request.data.get('query', {}) + xquery = dict(request.data.get('query') or {}) is_test = request.data.get('is_test', False) raise_exception = request.data.get('raise_exception', True) xquery['r_user'] = request.user.id xquery['r_dept'] = request.user.belong_dept.id if request.user.belong_dept else '' - can_cache = True - results = {} - results2 = {} - query = dt.default_param - if dt.sql_query: - if is_test: - query.update(dt.test_param) - else: - query.update(xquery) - try: - sql_f_ = check_sql_safe(dt.sql_query.format(**query)) - except KeyError as e: - raise ParseError(f'需指定查询参数_{str(e)}') - sql_f_strip = sql_f_.strip(';') - sql_f_l = sql_f_strip.split(';') + full_sql = render_dataset_sql(dt, xquery, is_test=is_test) + hash_k = None + if full_sql: + sql_f_strip = full_sql.strip(';') hash_k = hash(sql_f_strip) hash_v = cache.get(hash_k, None) if hash_v: return Response(hash_v) - # 多线程运行并返回字典结果 - with concurrent.futures.ThreadPoolExecutor(max_workers=6) as executor: - fun_ps = [] - for ind, val in enumerate(sql_f_l): - fun_ps.append((f'ds{ind}', execute_raw_sql, val)) - # 生成执行函数 - futures = {executor.submit(i[1], i[2]): i for i in fun_ps} - for future in concurrent.futures.as_completed(futures): - name, *_, sql_f = futures[future] # 获取对应的键 - try: - res = future.result() - results[name], results2[name] = format_sqldata( - res[0], res[1]) - except Exception as e: - myLogger.error(f'bi查询异常:{str(e)}-{dt.code}--{sql_f}') - if raise_exception: - raise ParseError(f'查询异常:{str(e)}') - else: - results[name] = 'error: ' + str(e) - can_cache = False - rdata['data'] = results - rdata['data2'] = results2 - if rdata['echart_options'] and not rdata['echart_options'].startswith('function'): - for key in results: - if isinstance(results[key], str): - raise ParseError(results[key]) - rdata['echart_options'] = format_json_with_placeholders( - rdata['echart_options'], **results) - if results and can_cache: + response_data, can_cache = execute_rendered_dataset( + dt, full_sql, raise_exception=raise_exception + ) + rdata.update(response_data) + if response_data['data'] and can_cache and hash_k is not None: cache.set(hash_k, rdata, dt.cache_seconds) return Response(rdata) diff --git a/docs/mcp.md b/docs/mcp.md new file mode 100644 index 00000000..ebde4537 --- /dev/null +++ b/docs/mcp.md @@ -0,0 +1,49 @@ +# Factory MCP 服务 + +Factory MCP 是仓库顶层的独立服务,使用官方 Python SDK v2,通过 Streamable HTTP 暴露 Agent 工具。当前已提供基础工具和第一批 Dataset 领域工具。 + +## 启动 + +在项目根目录使用项目虚拟环境启动独立进程: + +```powershell +.venv\Scripts\python.exe -m mcp_server +``` + +默认监听 `127.0.0.1:2260`,MCP 端点为 `/mcp`。客户端必须在每次请求中携带 Factory access token: + +```text +Authorization: Bearer +``` + +## 配置 + +生产环境在本机已忽略的 `config/conf.py` 中覆盖以下配置: + +- `MCP_HOST`:监听地址。 +- `MCP_PORT`:监听端口。 +- `MCP_PATH`:Streamable HTTP 路径。 +- `MCP_ALLOWED_HOSTS`:允许的 HTTP Host,支持 `hostname:*` 端口通配形式。 +- `MCP_ALLOWED_ORIGINS`:允许的浏览器 Origin;非浏览器客户端通常不发送 Origin。 +- `MCP_MAX_REQUEST_BODY_SIZE`:单个 MCP 请求体上限,默认 1 MiB。 +- `MCP_MAX_RESULT_BYTES`:单次领域工具结果上限,默认 512 KiB。 + +生产部署必须明确配置实际域名的 Host 白名单,不应直接复用 Django 当前的宽泛 `ALLOWED_HOSTS`。 + +## 基础工具 + +- `factory_server_info`:返回系统版本、MCP 协议版本和认证方式。 +- `factory_whoami`:返回当前 JWT 对应的 Factory 用户。 +- `search_datasets`:按名称、code 或描述搜索启用的数据集,不返回 SQL 配置。 +- `execute_dataset`:按 code 执行数据集,需要当前用户具有 `dataset.exec` 权限。 +- `search_wprs`:按编号、物料、批次、状态和当前位置搜索 WPR,只返回摘要。 +- `get_wpr`:按 ID、内部编号或对外编号读取 WPR 详情、缺陷和业务数据。 + +WPR 当前沿用既有 API 的读取边界:有效登录用户可读,且该 ViewSet 未启用部门数据过滤。MCP 不开放修改编号、分配对外编号或更新预处理信息等写操作。 + +- `search_batch_stats`:按批次、直通大批、起始物料和版本搜索批次统计摘要。 +- `get_batch_stat`:读取指定批次版本的完整统计数据,可附带直接拆批/合批关系。 + +BatchSt 同样沿用既有 API 的 `get: *` 读取边界,不提供创建、重算或修改工具。完整统计结果仍受 `MCP_MAX_RESULT_BYTES` 限制。 + +新增领域工具时必须从 MCP 请求身份获取用户,并复用 Factory 的权限码和数据范围过滤;不得直接使用固定管理员身份查询 ORM。 diff --git a/mcp_server/__init__.py b/mcp_server/__init__.py new file mode 100644 index 00000000..79fa5d7e --- /dev/null +++ b/mcp_server/__init__.py @@ -0,0 +1 @@ +"""Factory MCP v2 integration.""" diff --git a/mcp_server/__main__.py b/mcp_server/__main__.py new file mode 100644 index 00000000..1db0e253 --- /dev/null +++ b/mcp_server/__main__.py @@ -0,0 +1,24 @@ +import os + +import django +import uvicorn + + +def main() -> None: + os.environ.setdefault("DJANGO_SETTINGS_MODULE", "server.settings") + django.setup() + + from mcp_server.server import application + + from django.conf import settings + + uvicorn.run( + application, + host=settings.MCP_HOST, + port=settings.MCP_PORT, + log_level="info", + ) + + +if __name__ == "__main__": + main() diff --git a/mcp_server/auth.py b/mcp_server/auth.py new file mode 100644 index 00000000..180c1179 --- /dev/null +++ b/mcp_server/auth.py @@ -0,0 +1,66 @@ +from asgiref.sync import sync_to_async +from rest_framework.exceptions import APIException +from rest_framework_simplejwt.authentication import JWTAuthentication +from rest_framework_simplejwt.exceptions import TokenError + +from mcp.server.auth.provider import AccessToken +from mcp.server.auth.middleware.bearer_auth import RequireAuthMiddleware +from starlette.types import ASGIApp, Receive, Scope, Send + + +def verify_factory_jwt(token: str) -> AccessToken | None: + """验证 Factory access token,并生成 MCP 的逐请求身份信息。""" + authentication = JWTAuthentication() + try: + validated_token = authentication.get_validated_token(token) + user = authentication.get_user(validated_token) + except (APIException, TokenError): + return None + + return AccessToken( + token=token, + client_id="factory-mcp", + scopes=["factory:user"], + expires_at=validated_token.get("exp"), + subject=str(user.pk), + claims={ + "factory_user": { + "id": str(user.pk), + "username": user.get_username(), + "name": user.name, + "is_superuser": user.is_superuser, + }, + }, + ) + + +class FactoryJWTVerifier: + """让 MCP SDK 复用 Factory SimpleJWT 的验证规则。""" + + async def verify_token(self, token: str) -> AccessToken | None: + return await sync_to_async( + verify_factory_jwt, + thread_sensitive=True, + )(token) + + +class RequireFactoryJWTMiddleware: + """仅保护 HTTP 请求,并把 ASGI lifespan 原样交给 MCP SDK。""" + + def __init__(self, app: ASGIApp): + self.app = app + self.protected_app = RequireAuthMiddleware( + app, + required_scopes=[], + ) + + async def __call__( + self, + scope: Scope, + receive: Receive, + send: Send, + ) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return + await self.protected_app(scope, receive, send) diff --git a/mcp_server/context.py b/mcp_server/context.py new file mode 100644 index 00000000..087ddb48 --- /dev/null +++ b/mcp_server/context.py @@ -0,0 +1,32 @@ +from typing import Any + +from django.contrib.auth import get_user_model +from mcp.server.auth.middleware.auth_context import get_access_token + +from apps.utils.permission import has_perm + + +def authenticated_user_claims() -> dict[str, Any]: + """返回当前请求中经过 Factory JWT 校验的用户摘要。""" + access_token = get_access_token() + claims = access_token.claims if access_token else None + user = claims.get("factory_user") if claims else None + if not isinstance(user, dict): + raise RuntimeError("当前 MCP 请求缺少有效的 Factory 用户身份") + return user + + +def authenticated_factory_user(): + """加载当前 JWT 对应的 Django 用户,供权限和数据范围逻辑复用。""" + claims = authenticated_user_claims() + try: + return get_user_model().objects.select_related("belong_dept").get( + pk=claims["id"] + ) + except (KeyError, get_user_model().DoesNotExist) as exc: + raise RuntimeError("当前 JWT 对应的 Factory 用户不存在") from exc + + +def require_permission(user, permission_code: str) -> None: + if not has_perm(user, [permission_code]): + raise PermissionError(f"当前用户缺少权限:{permission_code}") diff --git a/mcp_server/server.py b/mcp_server/server.py new file mode 100644 index 00000000..e9b8a8cc --- /dev/null +++ b/mcp_server/server.py @@ -0,0 +1,89 @@ +from collections.abc import Sequence +from typing import Any + +from django.conf import settings +from mcp.server import MCPServer +from mcp.server.auth.middleware.auth_context import AuthContextMiddleware +from mcp.server.auth.middleware.bearer_auth import ( + BearerAuthBackend, +) +from mcp.server.auth.provider import TokenVerifier +from mcp.server.transport_security import TransportSecuritySettings +from starlette.middleware.authentication import AuthenticationMiddleware +from starlette.types import ASGIApp + +from mcp_server.auth import ( + FactoryJWTVerifier, + RequireFactoryJWTMiddleware, +) +from mcp_server.context import authenticated_user_claims +from mcp_server.tools.batch_stats import register_batch_stat_tools +from mcp_server.tools.datasets import register_dataset_tools +from mcp_server.tools.wprs import register_wpr_tools + + +PROTOCOL_REVISION = "2026-07-28" + +mcp = MCPServer( + name="factory", + title="Factory MCP", + description="Factory 面向 Agent 的受控业务能力入口。", + instructions="所有工具均使用当前请求携带的 Factory JWT 身份执行。", + version=settings.SYS_VERSION, +) + + +@mcp.tool() +def factory_server_info() -> dict[str, Any]: + """返回 Factory MCP 服务版本及协议基础信息。""" + return { + "name": "factory", + "system_version": settings.SYS_VERSION, + "protocol_revision": PROTOCOL_REVISION, + "authentication": "factory_jwt", + "domain_tools_ready": True, + } + + +@mcp.tool() +def factory_whoami() -> dict[str, Any]: + """返回当前 Factory JWT 对应的用户身份。""" + return authenticated_user_claims() + + +register_dataset_tools(mcp) +register_wpr_tools(mcp) +register_batch_stat_tools(mcp) + + +def create_app( + *, + allowed_hosts: Sequence[str] | None = None, + allowed_origins: Sequence[str] | None = None, + token_verifier: TokenVerifier | None = None, +) -> ASGIApp: + """创建仅接受 Factory JWT 的 MCP v2 Streamable HTTP 应用。""" + transport_security = TransportSecuritySettings( + enable_dns_rebinding_protection=True, + allowed_hosts=list(settings.MCP_ALLOWED_HOSTS if allowed_hosts is None else allowed_hosts), + allowed_origins=list(settings.MCP_ALLOWED_ORIGINS if allowed_origins is None else allowed_origins), + ) + app: ASGIApp = mcp.streamable_http_app( + streamable_http_path=settings.MCP_PATH, + json_response=True, + max_request_body_size=settings.MCP_MAX_REQUEST_BODY_SIZE, + transport_security=transport_security, + host=settings.MCP_HOST, + ) + + # 包装顺序保证先解析 Bearer JWT,再写入 MCP 请求上下文,最后强制认证。 + app = RequireFactoryJWTMiddleware(app) + app = AuthContextMiddleware(app) + app = AuthenticationMiddleware( + app, + backend=BearerAuthBackend(token_verifier or FactoryJWTVerifier()), + ) + return app + + +application = create_app() diff --git a/mcp_server/test_batch_stats.py b/mcp_server/test_batch_stats.py new file mode 100644 index 00000000..7d606638 --- /dev/null +++ b/mcp_server/test_batch_stats.py @@ -0,0 +1,105 @@ +from datetime import datetime +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from django.test import SimpleTestCase + +from apps.wpm.models import BatchSt +from mcp_server.tools.batch_stats import get_batch_stat, search_batch_stats + + +def batch_stat(**overrides): + values = { + "id": "500", + "batch": "BATCH-001", + "version": 1, + "zt_batch": "ZT-001", + "first_time": datetime(2026, 8, 1, 8, 0), + "last_time": datetime(2026, 8, 2, 8, 0), + "material_start": SimpleNamespace( + id="100", + name="原料", + model="M-1", + specification="S-1", + ), + "data": {"output": {"count": 10}, "quality": {"ok": 9}}, + "update_time": datetime(2026, 8, 10, 8, 0), + } + values.update(overrides) + return SimpleNamespace(**values) + + +class BatchStatToolTests(SimpleTestCase): + @patch("mcp_server.tools.batch_stats._base_queryset") + @patch("mcp_server.tools.batch_stats.authenticated_factory_user") + def test_search_returns_summary_without_full_data( + self, + _user_mock, + queryset_mock, + ): + queryset = MagicMock() + queryset.filter.return_value = queryset + queryset.order_by.return_value = queryset + queryset.__getitem__.return_value = [batch_stat()] + queryset_mock.return_value = queryset + + result = search_batch_stats(query="BATCH", limit=10) + + self.assertEqual(result["items"][0]["batch"], "BATCH-001") + self.assertEqual(result["items"][0]["data_keys"], ["output", "quality"]) + self.assertNotIn("data", result["items"][0]) + + @patch("mcp_server.tools.batch_stats.BatchLog.objects.filter") + @patch("mcp_server.tools.batch_stats._base_queryset") + @patch("mcp_server.tools.batch_stats.authenticated_factory_user") + def test_get_returns_data_and_direct_relations( + self, + _user_mock, + queryset_mock, + relation_filter_mock, + ): + item = batch_stat() + queryset_mock.return_value.get.return_value = item + relation_filter_mock.return_value.select_related.return_value.values.return_value = [ + { + "id": "600", + "relation_type": "split", + "source_id": "500", + "source__batch": "BATCH-001", + "source__version": 1, + "target_id": "501", + "target__batch": "BATCH-001-1", + "target__version": 1, + "handover_id": "700", + "mlog_id": None, + } + ] + + result = get_batch_stat("BATCH-001") + + self.assertEqual(result["data"]["output"]["count"], 10) + self.assertEqual(result["relations"][0]["relation_type"], "split") + + @patch("mcp_server.tools.batch_stats.BatchLog.objects.filter") + @patch("mcp_server.tools.batch_stats._base_queryset") + @patch("mcp_server.tools.batch_stats.authenticated_factory_user") + def test_get_can_omit_relations( + self, + _user_mock, + queryset_mock, + relation_filter_mock, + ): + queryset_mock.return_value.get.return_value = batch_stat() + + result = get_batch_stat("BATCH-001", include_relations=False) + + self.assertNotIn("relations", result) + relation_filter_mock.assert_not_called() + + @patch("mcp_server.tools.batch_stats._base_queryset") + @patch("mcp_server.tools.batch_stats.authenticated_factory_user") + def test_get_reports_missing_batch(self, _user_mock, queryset_mock): + queryset_mock.return_value.get.side_effect = BatchSt.DoesNotExist + + with self.assertRaisesRegex(ValueError, "未找到批次统计"): + get_batch_stat("missing") diff --git a/mcp_server/test_datasets.py b/mcp_server/test_datasets.py new file mode 100644 index 00000000..7d05cc77 --- /dev/null +++ b/mcp_server/test_datasets.py @@ -0,0 +1,96 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from django.test import SimpleTestCase, override_settings + +from mcp_server.tools.datasets import execute_dataset, search_datasets + + +class DatasetToolTests(SimpleTestCase): + @patch("mcp_server.tools.datasets.authenticated_factory_user") + @patch("mcp_server.tools.datasets.Dataset.objects.filter") + def test_search_returns_safe_catalog_fields(self, filter_mock, _user_mock): + queryset = MagicMock() + filter_mock.return_value = queryset + queryset.order_by.return_value.values.return_value.__getitem__.return_value = [ + { + "code": "daily_output", + "name": "日产量", + "description": "按日统计产量", + "default_param": {"day": "2026-08-10"}, + "test_param": {}, + } + ] + + result = search_datasets(limit=10) + + self.assertEqual(result["items"][0]["code"], "daily_output") + self.assertNotIn("sql_query", result["items"][0]) + + @patch("mcp_server.tools.datasets.cache") + @patch("mcp_server.tools.datasets.execute_rendered_dataset") + @patch("mcp_server.tools.datasets.render_dataset_sql") + @patch("mcp_server.tools.datasets.require_permission") + @patch("mcp_server.tools.datasets.authenticated_factory_user") + @patch("mcp_server.tools.datasets.Dataset.objects.get") + def test_execute_reuses_identity_permission_and_service( + self, + get_mock, + user_mock, + permission_mock, + render_mock, + execute_mock, + cache_mock, + ): + item = SimpleNamespace( + code="daily_output", + name="日产量", + description="按日统计产量", + cache_seconds=10, + ) + user = SimpleNamespace(id=42, belong_dept_id=7) + get_mock.return_value = item + user_mock.return_value = user + render_mock.return_value = "select 1" + cache_mock.get.return_value = None + execute_mock.return_value = ( + {"data": {"ds0": [{"count": 1}]}, "data2": {}}, + True, + ) + + result = execute_dataset("daily_output", {"day": "2026-08-10"}) + + permission_mock.assert_called_once_with(user, "dataset.exec") + render_query = render_mock.call_args.args[1] + self.assertEqual(render_query["r_user"], 42) + self.assertEqual(render_query["r_dept"], 7) + self.assertEqual(result["data"]["ds0"][0]["count"], 1) + self.assertNotIn("sql_query", result) + + @override_settings(MCP_MAX_RESULT_BYTES=1) + @patch("mcp_server.tools.datasets.cache") + @patch("mcp_server.tools.datasets.execute_rendered_dataset") + @patch("mcp_server.tools.datasets.render_dataset_sql", return_value="") + @patch("mcp_server.tools.datasets.require_permission") + @patch("mcp_server.tools.datasets.authenticated_factory_user") + @patch("mcp_server.tools.datasets.Dataset.objects.get") + def test_execute_rejects_oversized_results( + self, + get_mock, + user_mock, + _permission_mock, + _render_mock, + execute_mock, + _cache_mock, + ): + get_mock.return_value = SimpleNamespace( + code="daily_output", + name="日产量", + description="", + cache_seconds=0, + ) + user_mock.return_value = SimpleNamespace(id=42, belong_dept_id=None) + execute_mock.return_value = ({"data": {"ds0": [1]}, "data2": {}}, True) + + with self.assertRaisesRegex(RuntimeError, "超过 MCP 响应上限"): + execute_dataset("daily_output") diff --git a/mcp_server/test_wprs.py b/mcp_server/test_wprs.py new file mode 100644 index 00000000..ed33b313 --- /dev/null +++ b/mcp_server/test_wprs.py @@ -0,0 +1,103 @@ +from datetime import datetime +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from django.test import SimpleTestCase + +from mcp_server.tools.wprs import get_wpr, search_wprs + + +def material(material_id="100", name="成品"): + return SimpleNamespace( + id=material_id, + name=name, + model="M-1", + specification="S-1", + ) + + +def wpr(**overrides): + values = { + "id": "200", + "number": "WPR-001", + "number_out": "OUT-001", + "version": 1, + "state": 10, + "get_state_display": lambda: "正常", + "material": material(), + "material_start": material("101", "原料"), + "wm_id": "300", + "wm": SimpleNamespace(batch="WP-001"), + "mb_id": None, + "mb": None, + "wpr_from_id": None, + "wpr_from": None, + "oinfo": {"test": "ok"}, + "data": {"route": []}, + "pre_info": {"tooling": "T-1"}, + "create_time": datetime(2026, 8, 1, 8, 0), + "update_time": datetime(2026, 8, 10, 8, 0), + } + values.update(overrides) + return SimpleNamespace(**values) + + +class WprToolTests(SimpleTestCase): + @patch("mcp_server.tools.wprs._base_queryset") + @patch("mcp_server.tools.wprs.authenticated_factory_user") + def test_search_returns_read_only_summary(self, _user_mock, queryset_mock): + queryset = MagicMock() + queryset.filter.return_value = queryset + queryset.distinct.return_value = queryset + queryset.order_by.return_value = queryset + queryset.__getitem__.return_value = [wpr()] + queryset_mock.return_value = queryset + + result = search_wprs(query="WPR", location="workshop", limit=10) + + self.assertEqual(result["items"][0]["number"], "WPR-001") + self.assertEqual(result["items"][0]["workshop_batch"], "WP-001") + self.assertNotIn("data", result["items"][0]) + self.assertNotIn("pre_info", result["items"][0]) + + @patch("mcp_server.tools.wprs.WprDefect.objects.filter") + @patch("mcp_server.tools.wprs._base_queryset") + @patch("mcp_server.tools.wprs.authenticated_factory_user") + def test_get_returns_business_detail( + self, + _user_mock, + queryset_mock, + defect_filter_mock, + ): + item = wpr() + queryset = MagicMock() + queryset.filter.return_value.order_by.return_value.first.return_value = item + queryset_mock.return_value = queryset + defect_filter_mock.return_value.select_related.return_value.values.return_value = [ + { + "defect_id": "400", + "defect__name": "划伤", + "is_main": True, + } + ] + + result = get_wpr("WPR-001") + + self.assertEqual(result["material"]["name"], "成品") + self.assertEqual(result["defects"][0]["defect__name"], "划伤") + self.assertEqual(result["pre_info"]["tooling"], "T-1") + + @patch("mcp_server.tools.wprs._base_queryset") + @patch("mcp_server.tools.wprs.authenticated_factory_user") + def test_get_reports_missing_wpr(self, _user_mock, queryset_mock): + queryset = MagicMock() + queryset.filter.return_value.order_by.return_value.first.return_value = None + queryset_mock.return_value = queryset + + with self.assertRaisesRegex(ValueError, "未找到 WPR"): + get_wpr("missing") + + @patch("mcp_server.tools.wprs.authenticated_factory_user") + def test_search_rejects_unknown_location(self, _user_mock): + with self.assertRaisesRegex(ValueError, "不支持的 WPR 位置"): + search_wprs(location="invalid") diff --git a/mcp_server/tests.py b/mcp_server/tests.py new file mode 100644 index 00000000..8e4fa8d6 --- /dev/null +++ b/mcp_server/tests.py @@ -0,0 +1,164 @@ +import asyncio +from types import SimpleNamespace +from unittest.mock import patch + +from django.test import SimpleTestCase +from mcp.server.auth.middleware.auth_context import auth_context_var +from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser +from mcp.server.auth.provider import AccessToken +from rest_framework.exceptions import AuthenticationFailed +from starlette.testclient import TestClient + +from mcp_server.auth import verify_factory_jwt +from mcp_server.server import ( + PROTOCOL_REVISION, + create_app, + factory_server_info, + factory_whoami, + mcp, +) + + +class FactoryJWTVerifierTests(SimpleTestCase): + def test_factory_access_token_is_accepted(self): + user = SimpleNamespace( + pk=42, + name="MCP用户", + is_superuser=False, + get_username=lambda: "mcp-user", + ) + authentication = patch("mcp_server.auth.JWTAuthentication").start() + self.addCleanup(patch.stopall) + authentication.return_value.get_validated_token.return_value = { + "exp": 1234567890, + } + authentication.return_value.get_user.return_value = user + + access_token = verify_factory_jwt("access-token") + + self.assertIsNotNone(access_token) + self.assertEqual(access_token.subject, str(user.pk)) + self.assertEqual( + access_token.claims["factory_user"]["username"], + user.get_username(), + ) + + def test_invalid_token_is_rejected(self): + with patch("mcp_server.auth.JWTAuthentication") as authentication: + authentication.return_value.get_validated_token.side_effect = AuthenticationFailed("invalid token") + + self.assertIsNone(verify_factory_jwt("invalid-token")) + + +class FactoryMCPServerTests(SimpleTestCase): + def test_base_tools_are_registered(self): + tools = asyncio.run(mcp.list_tools()) + names = {tool.name for tool in tools} + + self.assertEqual( + names, + { + "execute_dataset", + "factory_server_info", + "factory_whoami", + "get_batch_stat", + "get_wpr", + "search_batch_stats", + "search_datasets", + "search_wprs", + }, + ) + + def test_server_info_targets_mcp_v2(self): + result = factory_server_info() + + self.assertEqual(result["protocol_revision"], PROTOCOL_REVISION) + self.assertTrue(result["domain_tools_ready"]) + + def test_whoami_uses_authenticated_request_context(self): + user = { + "id": "42", + "username": "agent-user", + "name": "Agent用户", + "is_superuser": False, + } + authenticated = AuthenticatedUser( + AccessToken( + token="test-token", + client_id="factory-mcp", + scopes=["factory:user"], + subject=user["id"], + claims={"factory_user": user}, + ) + ) + context_token = auth_context_var.set(authenticated) + try: + self.assertEqual(factory_whoami(), user) + finally: + auth_context_var.reset(context_token) + + def test_http_endpoint_requires_bearer_token(self): + app = create_app(allowed_hosts=["testserver"]) + + with TestClient(app) as client: + response = client.post("/mcp", json={}) + + self.assertEqual(response.status_code, 401) + self.assertEqual(response.json()["error"], "invalid_token") + + def test_mcp_v2_request_uses_bearer_identity(self): + user = { + "id": "42", + "username": "agent-user", + "name": "Agent用户", + "is_superuser": False, + } + + class TestTokenVerifier: + async def verify_token(self, token): + if token != "valid-token": + return None + return AccessToken( + token=token, + client_id="test-client", + scopes=["factory:user"], + subject=user["id"], + claims={"factory_user": user}, + ) + + app = create_app( + allowed_hosts=["testserver"], + token_verifier=TestTokenVerifier(), + ) + request = { + "jsonrpc": "2.0", + "id": 1, + "method": "tools/call", + "params": { + "name": "factory_whoami", + "arguments": {}, + "_meta": { + "io.modelcontextprotocol/protocolVersion": (PROTOCOL_REVISION), + "io.modelcontextprotocol/clientInfo": { + "name": "factory-tests", + "version": "1.0", + }, + "io.modelcontextprotocol/clientCapabilities": {}, + }, + }, + } + headers = { + "Authorization": "Bearer valid-token", + "MCP-Protocol-Version": PROTOCOL_REVISION, + "Mcp-Method": "tools/call", + "Mcp-Name": "factory_whoami", + } + + with TestClient(app) as client: + response = client.post("/mcp", json=request, headers=headers) + + self.assertEqual(response.status_code, 200) + self.assertEqual( + response.json()["result"]["structuredContent"], + user, + ) diff --git a/mcp_server/tools/__init__.py b/mcp_server/tools/__init__.py new file mode 100644 index 00000000..f7330c3a --- /dev/null +++ b/mcp_server/tools/__init__.py @@ -0,0 +1 @@ +"""Factory MCP 领域工具,按业务域拆分并在 server 中显式注册。""" diff --git a/mcp_server/tools/batch_stats.py b/mcp_server/tools/batch_stats.py new file mode 100644 index 00000000..ed1fd251 --- /dev/null +++ b/mcp_server/tools/batch_stats.py @@ -0,0 +1,104 @@ +from typing import Any + +from django.db.models import Q + +from apps.wpm.models import BatchLog, BatchSt +from mcp_server.context import authenticated_factory_user +from mcp_server.tools.common import json_safe_result, validate_result_size + + +def _base_queryset(): + return BatchSt.objects.select_related("material_start") + + +def _batch_summary(batch_stat: BatchSt) -> dict[str, Any]: + material = batch_stat.material_start + return { + "id": str(batch_stat.id), + "batch": batch_stat.batch, + "version": batch_stat.version, + "zt_batch": batch_stat.zt_batch, + "first_time": batch_stat.first_time, + "last_time": batch_stat.last_time, + "material_start": ( + { + "id": str(material.id), + "name": material.name, + "model": material.model, + "specification": material.specification, + } + if material + else None + ), + "data_keys": sorted((batch_stat.data or {}).keys()), + "update_time": batch_stat.update_time, + } + + +def search_batch_stats( + query: str = "", + zt_batch: str = "", + material_id: str | None = None, + version: int | None = 1, + limit: int = 20, +) -> dict[str, Any]: + """搜索批次统计;摘要仅返回数据分组名称,不返回完整统计数据。""" + authenticated_factory_user() + safe_limit = max(1, min(limit, 100)) + queryset = _base_queryset() + if query.strip(): + queryset = queryset.filter(batch__icontains=query.strip()) + if zt_batch.strip(): + queryset = queryset.filter(zt_batch=zt_batch.strip()) + if material_id: + queryset = queryset.filter(material_start_id=material_id) + if version is not None: + queryset = queryset.filter(version=version) + items = [ + _batch_summary(item) + for item in queryset.order_by("batch", "version")[:safe_limit] + ] + result = json_safe_result({"items": items, "limit": safe_limit}) + validate_result_size(result) + return result + + +def get_batch_stat( + batch: str, + version: int = 1, + include_relations: bool = True, +) -> dict[str, Any]: + """读取指定批次版本的完整统计数据,并可附带直接拆合批关系。""" + authenticated_factory_user() + try: + batch_stat = _base_queryset().get(batch=batch, version=version) + except BatchSt.DoesNotExist as exc: + raise ValueError(f"未找到批次统计:{batch} v{version}") from exc + + result = _batch_summary(batch_stat) + result["data"] = batch_stat.data + if include_relations: + result["relations"] = list( + BatchLog.objects.filter(Q(source=batch_stat) | Q(target=batch_stat)) + .select_related("source", "target") + .values( + "id", + "relation_type", + "source_id", + "source__batch", + "source__version", + "target_id", + "target__batch", + "target__version", + "handover_id", + "mlog_id", + ) + ) + result = json_safe_result(result) + validate_result_size(result) + return result + + +def register_batch_stat_tools(server) -> None: + server.tool()(search_batch_stats) + server.tool()(get_batch_stat) diff --git a/mcp_server/tools/common.py b/mcp_server/tools/common.py new file mode 100644 index 00000000..61aa872b --- /dev/null +++ b/mcp_server/tools/common.py @@ -0,0 +1,18 @@ +import json +from typing import Any + +from django.conf import settings + +from apps.utils.tools import MyJSONEncoder + + +def json_safe_result(value: Any) -> Any: + return json.loads(json.dumps(value, cls=MyJSONEncoder, ensure_ascii=False)) + + +def validate_result_size(value: Any) -> None: + encoded = json.dumps(value, ensure_ascii=False, separators=(",", ":")).encode() + if len(encoded) > settings.MCP_MAX_RESULT_BYTES: + raise RuntimeError( + "工具结果超过 MCP 响应上限,请缩小查询范围或增加筛选参数" + ) diff --git a/mcp_server/tools/datasets.py b/mcp_server/tools/datasets.py new file mode 100644 index 00000000..d78155d2 --- /dev/null +++ b/mcp_server/tools/datasets.py @@ -0,0 +1,79 @@ +import hashlib +from typing import Any + +from django.core.cache import cache +from django.db.models import Q + +from apps.bi.models import Dataset +from apps.bi.services import execute_rendered_dataset, render_dataset_sql +from mcp_server.context import authenticated_factory_user, require_permission +from mcp_server.tools.common import json_safe_result, validate_result_size + + +def search_datasets(query: str = "", limit: int = 20) -> dict[str, Any]: + """搜索可执行的数据集目录,不返回 SQL 等敏感配置。""" + authenticated_factory_user() + safe_limit = max(1, min(limit, 100)) + queryset = Dataset.objects.filter(enabled=True) + if query.strip(): + queryset = queryset.filter( + Q(name__icontains=query.strip()) + | Q(code__icontains=query.strip()) + | Q(description__icontains=query.strip()) + ) + rows = queryset.order_by("name", "code", "id").values( + "code", + "name", + "description", + "default_param", + "test_param", + )[:safe_limit] + return {"items": list(rows), "limit": safe_limit} + + +def execute_dataset( + code: str, + parameters: dict[str, Any] | None = None, +) -> dict[str, Any]: + """以当前 Factory 用户身份执行启用的数据集。需要 dataset.exec 权限。""" + user = authenticated_factory_user() + require_permission(user, "dataset.exec") + try: + dataset = Dataset.objects.get(code=code, enabled=True) + except Dataset.DoesNotExist as exc: + raise ValueError(f"未找到已启用的数据集:{code}") from exc + except Dataset.MultipleObjectsReturned as exc: + raise RuntimeError(f"数据集 code 不唯一,无法执行:{code}") from exc + + query = dict(parameters or {}) + query["r_user"] = user.id + query["r_dept"] = user.belong_dept_id or "" + full_sql = render_dataset_sql(dataset, query) + + cache_key = None + response_data = None + if full_sql and dataset.cache_seconds: + digest = hashlib.sha256(full_sql.strip(";").encode()).hexdigest() + cache_key = f"mcp:dataset:{digest}" + response_data = cache.get(cache_key) + + if response_data is None: + response_data, can_cache = execute_rendered_dataset(dataset, full_sql) + if cache_key and can_cache and response_data["data"]: + cache.set(cache_key, response_data, dataset.cache_seconds) + + result = json_safe_result( + { + "code": dataset.code, + "name": dataset.name, + "description": dataset.description, + **response_data, + } + ) + validate_result_size(result) + return result + + +def register_dataset_tools(server) -> None: + server.tool()(search_datasets) + server.tool()(execute_dataset) diff --git a/mcp_server/tools/wprs.py b/mcp_server/tools/wprs.py new file mode 100644 index 00000000..d8b1996a --- /dev/null +++ b/mcp_server/tools/wprs.py @@ -0,0 +1,139 @@ +from typing import Any, Literal + +from django.db.models import Q + +from apps.wpmw.models import Wpr, WprDefect +from mcp_server.context import authenticated_factory_user +from mcp_server.tools.common import json_safe_result, validate_result_size + + +WprLocation = Literal["all", "workshop", "warehouse", "unassigned"] + + +def _base_queryset(): + return Wpr.objects.select_related( + "material", + "material_start", + "wm", + "mb", + "wpr_from", + ) + + +def _wpr_summary(wpr: Wpr) -> dict[str, Any]: + material = wpr.material + return { + "id": str(wpr.id), + "number": wpr.number, + "number_out": wpr.number_out, + "version": wpr.version, + "state": wpr.state, + "state_name": wpr.get_state_display(), + "material": { + "id": str(material.id), + "name": material.name, + "model": material.model, + "specification": material.specification, + }, + "workshop_batch": wpr.wm.batch if wpr.wm_id else None, + "warehouse_batch": wpr.mb.batch if wpr.mb_id else None, + "create_time": wpr.create_time, + "update_time": wpr.update_time, + } + + +def search_wprs( + query: str = "", + state: int | None = None, + material_id: str | None = None, + batch: str = "", + location: WprLocation = "all", + limit: int = 20, +) -> dict[str, Any]: + """按编号、物料或批次搜索单件产品;仅提供只读摘要。""" + authenticated_factory_user() + if location not in {"all", "workshop", "warehouse", "unassigned"}: + raise ValueError(f"不支持的 WPR 位置:{location}") + safe_limit = max(1, min(limit, 100)) + queryset = _base_queryset() + if query.strip(): + keyword = query.strip() + queryset = queryset.filter( + Q(number__icontains=keyword) + | Q(number_out__icontains=keyword) + | Q(material__name__icontains=keyword) + | Q(material__model__icontains=keyword) + | Q(material__specification__icontains=keyword) + ) + if state is not None: + queryset = queryset.filter(state=state) + if material_id: + queryset = queryset.filter(material_id=material_id) + if batch.strip(): + queryset = queryset.filter( + Q(wm__batch__icontains=batch.strip()) + | Q(mb__batch__icontains=batch.strip()) + ) + if location == "workshop": + queryset = queryset.filter(wm__isnull=False) + elif location == "warehouse": + queryset = queryset.filter(mb__isnull=False) + elif location == "unassigned": + queryset = queryset.filter(wm__isnull=True, mb__isnull=True) + + items = [ + _wpr_summary(wpr) + for wpr in queryset.distinct().order_by("number", "create_time")[:safe_limit] + ] + result = json_safe_result({"items": items, "limit": safe_limit}) + validate_result_size(result) + return result + + +def get_wpr(identifier: str) -> dict[str, Any]: + """按 WPR ID、内部编号或对外编号读取单件详情。""" + authenticated_factory_user() + lookup = Q(number=identifier) | Q(number_out=identifier) + if identifier.isdigit(): + lookup |= Q(pk=identifier) + wpr = _base_queryset().filter(lookup).order_by("-version", "-update_time").first() + if wpr is None: + raise ValueError(f"未找到 WPR:{identifier}") + + result = _wpr_summary(wpr) + material_start = wpr.material_start + result.update( + { + "material_start": ( + { + "id": str(material_start.id), + "name": material_start.name, + "model": material_start.model, + "specification": material_start.specification, + } + if material_start + else None + ), + "wpr_from": ( + {"id": str(wpr.wpr_from.id), "number": wpr.wpr_from.number} + if wpr.wpr_from_id + else None + ), + "oinfo": wpr.oinfo, + "data": wpr.data, + "pre_info": wpr.pre_info, + "defects": list( + WprDefect.objects.filter(wpr=wpr) + .select_related("defect") + .values("defect_id", "defect__name", "is_main") + ), + } + ) + result = json_safe_result(result) + validate_result_size(result) + return result + + +def register_wpr_tools(server) -> None: + server.tool()(search_wprs) + server.tool()(get_wpr) diff --git a/requirements.txt b/requirements.txt index 0f5f390e..49d8cf82 100755 --- a/requirements.txt +++ b/requirements.txt @@ -10,6 +10,11 @@ django-cors-headers==4.9.0 djangorestframework-simplejwt==5.5.1 django-restql==0.15.2 +# ======================= +# Agent Integration +# ======================= +mcp==2.0.0 + # ======================= # Celery # ======================= diff --git a/server/settings.py b/server/settings.py index 1df8238b..934f788a 100755 --- a/server/settings.py +++ b/server/settings.py @@ -241,6 +241,20 @@ SIMPLE_JWT = { 'REFRESH_TOKEN_LIFETIME': timedelta(days=60), } +# MCP v2 服务配置。生产环境应在 config/conf.py 中覆盖监听地址和 Host/Origin 白名单。 +MCP_HOST = globals().get('MCP_HOST', '127.0.0.1') +MCP_PORT = globals().get('MCP_PORT', 2260) +MCP_PATH = globals().get('MCP_PATH', '/mcp') +MCP_ALLOWED_HOSTS = globals().get( + 'MCP_ALLOWED_HOSTS', + ['127.0.0.1', '127.0.0.1:*', 'localhost', 'localhost:*'], +) +MCP_ALLOWED_ORIGINS = globals().get('MCP_ALLOWED_ORIGINS', []) +MCP_MAX_REQUEST_BODY_SIZE = globals().get( + 'MCP_MAX_REQUEST_BODY_SIZE', 1024 * 1024 +) +MCP_MAX_RESULT_BYTES = globals().get('MCP_MAX_RESULT_BYTES', 512 * 1024) + # 跨域配置/可用nginx处理,无需引入corsheaders CORS_ORIGIN_ALLOW_ALL = True CORS_ALLOW_CREDENTIALS = True