This commit is contained in:
shijing 2026-08-11 16:02:57 +08:00
commit 7e318d9cfd
64 changed files with 3245 additions and 191 deletions

View File

@ -8,6 +8,8 @@
- [合批原料字段历史问题](project_material_ofrom_merge_bug.md)`material_ofrom` 不一致的既有排查结论。 - [合批原料字段历史问题](project_material_ofrom_merge_bug.md)`material_ofrom` 不一致的既有排查结论。
- [前端独立发版流程](reference_ehs_web_release.md):发布 `ehs_web` 时使用。 - [前端独立发版流程](reference_ehs_web_release.md):发布 `ehs_web` 时使用。
- [项目 Python 虚拟环境](reference_python_venv.md):运行 Django、pytest 和脚本时使用。 - [项目 Python 虚拟环境](reference_python_venv.md):运行 Django、pytest 和脚本时使用。
- [项目测试数据库](reference_test_database.md):测试与可切换的生产查询连接解耦,始终使用固定的 `test_ehs_develop`
- [两个 WebView 套壳 App](reference_wrapper_apps.md):修改 h5x 与扫码、返回键交互时使用。 - [两个 WebView 套壳 App](reference_wrapper_apps.md):修改 h5x 与扫码、返回键交互时使用。
- [前端验证时机](feedback_frontend_validation.md):日常修改先跑 check完整 build 留到 push 前执行。
这些文件记录的是长期约定或历史上下文。执行任务前应结合当前代码和数据重新验证,尤其不要把历史缺陷结论直接当成当前故障原因。 这些文件记录的是长期约定或历史上下文。执行任务前应结合当前代码和数据重新验证,尤其不要把历史缺陷结论直接当成当前故障原因。

View File

@ -0,0 +1,5 @@
# 前端验证时机
- 修改配套前端 `../ehs_web` 时,日常开发和中间验证优先运行项目已有的 `check`,不要每次修改后都运行完整 `build`
- 准备 push 前运行一次完整 `build`,用于发现生产构建阶段的问题。
- 若当前前端尚未配置 `check` 脚本,应先说明现状,不得把其他命令擅自当作 `check`

View File

@ -0,0 +1,15 @@
# 项目测试数据库
- 生产问题查询时,默认数据库连接可以根据工厂或环境切换到不同 IP 和业务库;这类连接只用于授权范围内的生产数据只读验证。
- 所有 Django、pytest 及其他自动化测试必须与当前生产查询连接解耦,始终使用固定的 `test_ehs_develop` 测试数据库连接。
- 不能仅依赖当前默认连接的 Django 自动 `test_<NAME>` 命名;即使当前业务库是 `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`

View File

@ -18,10 +18,35 @@ class DatasetCreateUpdateSerializer(CustomModelSerializer):
class DatasetSerializer(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: class Meta:
model = Dataset model = Dataset
fields = '__all__' 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 DatasetRecordSerializer(CustomModelSerializer):
class Meta: class Meta:
model = DatasetRecord model = DatasetRecord
@ -36,6 +61,14 @@ class DatasetRecordSerializer(CustomModelSerializer):
class DataExecSerializer(serializers.Serializer): class DataExecSerializer(serializers.Serializer):
query = serializers.JSONField( query = serializers.JSONField(
label="查询字典参数", required=False, allow_null=True) label="查询字典参数",
is_test = serializers.BooleanField(label='是否测试', default=False) help_text="按所选数据集 description/default_param 声明的业务参数填写",
raise_exception = serializers.BooleanField(label='是否直接报错', default=False) 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
)

View File

@ -1,11 +1,15 @@
from rest_framework.exceptions import ParseError import concurrent.futures
import json import json
from jinja2 import Template import logging
from rest_framework.exceptions import ParseError
from apps.bi.models import Dataset from apps.bi.models import Dataset
import concurrent
from apps.utils.sql import execute_raw_sql, format_sqldata from apps.utils.sql import execute_raw_sql, format_sqldata
from apps.utils.tools import MyJSONEncoder from apps.utils.tools import MyJSONEncoder
myLogger = logging.getLogger('log')
forbidden_keywords = ["UPDATE", "DELETE", "DROP", "TRUNCATE", "INSERT", "CREATE", "ALTER", "GRANT", "REVOKE", "EXEC", "EXECUTE"] 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 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}) 返回 (sql语句, { rda})
""" """
rdata = {} full_sql = render_dataset_sql(dt, xquery)
results = {} response_data, _ = execute_rendered_dataset(dt, full_sql)
results2 = {} return full_sql, response_data
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

105
apps/bi/test_services.py Normal file
View File

@ -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)

View File

@ -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)

View File

@ -3,17 +3,21 @@ from apps.utils.viewsets import CustomModelViewSet, CustomGenericViewSet
from rest_framework.decorators import action from rest_framework.decorators import action
from rest_framework.response import Response from rest_framework.response import Response
from apps.bi.models import Dataset, DatasetRecord 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 from django.apps import apps
import concurrent.futures
from django.core.cache import cache from django.core.cache import cache
from apps.utils.sql import execute_raw_sql, format_sqldata from apps.bi.services import execute_rendered_dataset, render_dataset_sql
from apps.bi.services import check_sql_safe, format_json_with_placeholders
from rest_framework.exceptions import ParseError from rest_framework.exceptions import ParseError
from rest_framework.generics import get_object_or_404 from rest_framework.generics import get_object_or_404
from apps.utils.mixins import ListModelMixin from apps.utils.mixins import ListModelMixin
import logging from drf_yasg import openapi
myLogger = logging.getLogger('log') from drf_yasg.utils import swagger_auto_schema
# Create your views here. # Create your views here.
@ -22,9 +26,54 @@ class DatasetViewSet(CustomModelViewSet):
serializer_class = DatasetSerializer serializer_class = DatasetSerializer
create_serializer_class = DatasetCreateUpdateSerializer create_serializer_class = DatasetCreateUpdateSerializer
update_serializer_class = DatasetCreateUpdateSerializer update_serializer_class = DatasetCreateUpdateSerializer
search_fields = ['name', 'code'] search_fields = ['name', 'code', 'description']
ordering = ['name', 'code', 'id'] 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): def get_object(self):
""" """
Returns the object the view is displaying. Returns the object the view is displaying.
@ -57,6 +106,18 @@ class DatasetViewSet(CustomModelViewSet):
return obj 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=[]) @action(methods=['post'], detail=True, perms_map={'post': 'dataset.exec'}, serializer_class=DataExecSerializer, cache_seconds=0, logging_methods=[])
def exec(self, request, pk=None): def exec(self, request, pk=None):
"""执行sql查询 """执行sql查询
@ -67,59 +128,24 @@ class DatasetViewSet(CustomModelViewSet):
if not dt.enabled: if not dt.enabled:
raise ParseError(f'{dt.name}-该查询未启用') raise ParseError(f'{dt.name}-该查询未启用')
rdata = DatasetSerializer(instance=dt).data 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) is_test = request.data.get('is_test', False)
raise_exception = request.data.get('raise_exception', True) raise_exception = request.data.get('raise_exception', True)
xquery['r_user'] = request.user.id xquery['r_user'] = request.user.id
xquery['r_dept'] = request.user.belong_dept.id if request.user.belong_dept else '' xquery['r_dept'] = request.user.belong_dept.id if request.user.belong_dept else ''
can_cache = True full_sql = render_dataset_sql(dt, xquery, is_test=is_test)
results = {} hash_k = None
results2 = {} if full_sql:
query = dt.default_param sql_f_strip = full_sql.strip(';')
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(';')
hash_k = hash(sql_f_strip) hash_k = hash(sql_f_strip)
hash_v = cache.get(hash_k, None) hash_v = cache.get(hash_k, None)
if hash_v: if hash_v:
return Response(hash_v) return Response(hash_v)
# 多线程运行并返回字典结果 response_data, can_cache = execute_rendered_dataset(
with concurrent.futures.ThreadPoolExecutor(max_workers=6) as executor: dt, full_sql, raise_exception=raise_exception
fun_ps = [] )
for ind, val in enumerate(sql_f_l): rdata.update(response_data)
fun_ps.append((f'ds{ind}', execute_raw_sql, val)) if response_data['data'] and can_cache and hash_k is not None:
# 生成执行函数
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:
cache.set(hash_k, rdata, dt.cache_seconds) cache.set(hash_k, rdata, dt.cache_seconds)
return Response(rdata) return Response(rdata)

View File

@ -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])

View File

@ -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, from apps.develop.views import (BackupDatabase, BackupMedia, ReloadClientGit,
ReloadServerGit, ReloadServerOnly, TestViewSet, CorrectViewSet, testScanHtml, ServerTime) ReloadServerGit, ReloadServerOnly, TestViewSet, CorrectViewSet, testScanHtml, ServerTime)
from rest_framework.routers import DefaultRouter from rest_framework.routers import DefaultRouter
@ -6,9 +7,11 @@ from rest_framework.routers import DefaultRouter
API_BASE_URL = 'api/develop/' API_BASE_URL = 'api/develop/'
HTML_BASE_URL = 'dhtml/develop/' HTML_BASE_URL = 'dhtml/develop/'
router = DefaultRouter() router = DefaultRouter()
router.register('test', TestViewSet, basename='api_test')
router.register('correct', CorrectViewSet, basename='correct') router.register('correct', CorrectViewSet, basename='correct')
if settings.DEBUG:
router.register('test', TestViewSet, basename='api_test')
urlpatterns = [ urlpatterns = [
path(API_BASE_URL + 'reload_server_git/', ReloadServerGit.as_view()), path(API_BASE_URL + 'reload_server_git/', ReloadServerGit.as_view()),
# path(API_BASE_URL + 'reload_web_git/', ReloadClientGit.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 + 'backup_media/', BackupMedia.as_view()),
path(API_BASE_URL + 'server_time/', ServerTime.as_view()), path(API_BASE_URL + 'server_time/', ServerTime.as_view()),
path(API_BASE_URL, include(router.urls)), path(API_BASE_URL, include(router.urls)),
path(HTML_BASE_URL + "testscan/", testScanHtml)
] ]
if settings.DEBUG:
urlpatterns.append(path(HTML_BASE_URL + "testscan/", testScanHtml))

View File

@ -2,7 +2,7 @@
from rest_framework.views import APIView from rest_framework.views import APIView
from rest_framework.exceptions import ParseError 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.response import Response
from rest_framework.serializers import Serializer from rest_framework.serializers import Serializer
from rest_framework.decorators import action from rest_framework.decorators import action
@ -40,11 +40,7 @@ from datetime import datetime
# Create your views here. # Create your views here.
class ServerTime(APIView): class ServerTime(APIView):
permission_classes = [IsAdminUser]
def get_permissions(self):
if self.request.method == 'GET':
return [AllowAny()]
return [IsAdminUser()]
@swagger_auto_schema(responses={200: ServerTimeSerializer}) @swagger_auto_schema(responses={200: ServerTimeSerializer})
def get(self, request): 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( completed = subprocess.run(
["sudo", "-S", "sh", "-c", command], # 添加 -S 参数 ["sudo", "-S", "date", "-s", server_time],
input=SD_PWD + "\n", # 注意要在密码后加换行符 input=SD_PWD + "\n", # 注意要在密码后加换行符
capture_output=True, capture_output=True,
text=True text=True
@ -269,10 +269,9 @@ class CorrectViewSet(CustomGenericViewSet):
class TestViewSet(CustomGenericViewSet): class TestViewSet(CustomGenericViewSet):
perms_map = {} perms_map = {}
authentication_classes = () permission_classes = [IsAdminUser]
permission_classes = ()
@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): def send_sms(self, request, pk=None):
"""发送短信测试 """发送短信测试
@ -565,7 +564,7 @@ class TestViewSet(CustomGenericViewSet):
# correct_card_time() # correct_card_time()
# return Response() # return Response()
@action(methods=['post'], detail=False, serializer_class=Serializer, permission_classes=[]) @action(methods=['post'], detail=False, serializer_class=Serializer)
@transaction.atomic @transaction.atomic
def correct_data(self, request, pk=None): def correct_data(self, request, pk=None):
"""修正数据 """修正数据
@ -674,7 +673,7 @@ class TestViewSet(CustomGenericViewSet):
Ticket.objects.get_queryset(all=True).delete() Ticket.objects.get_queryset(all=True).delete()
return Response() 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): def test_cal(self, request, pk=None):
from apps.wpm.tasks import cal_exp_duration_sec from apps.wpm.tasks import cal_exp_duration_sec
cal_exp_duration_sec('3397169058570170368') cal_exp_duration_sec('3397169058570170368')
@ -710,4 +709,4 @@ html_str = """
</html> </html>
""" """
def testScanHtml(request): def testScanHtml(request):
return HttpResponse(html_str) return HttpResponse(html_str)

View File

@ -71,6 +71,8 @@ class ExamViewSet(CustomModelViewSet):
def get_queryset(self): def get_queryset(self):
qs = super().get_queryset() qs = super().get_queryset()
if getattr(self, 'swagger_fake_view', False):
return qs
if has_perm(self.request.user, ["exam.view"]): if has_perm(self.request.user, ["exam.view"]):
return qs return qs
user:User = self.request.user user:User = self.request.user
@ -142,6 +144,8 @@ class ExamRecordViewSet(ListModelMixin, DestroyModelMixin, RetrieveModelMixin, C
def get_queryset(self): def get_queryset(self):
qs = super().get_queryset() qs = super().get_queryset()
if getattr(self, 'swagger_fake_view', False):
return qs
if has_perm(self.request.user, ["examrecord.view"]): if has_perm(self.request.user, ["examrecord.view"]):
return qs return qs
return qs.filter(create_by=self.request.user) return qs.filter(create_by=self.request.user)
@ -207,6 +211,8 @@ class TrainRecordViewSet(CustomModelViewSet):
def get_queryset(self): def get_queryset(self):
qs = super().get_queryset() qs = super().get_queryset()
if getattr(self, 'swagger_fake_view', False):
return qs
if has_perm(self.request.user, ["train.view"]): if has_perm(self.request.user, ["train.view"]):
return qs return qs
return qs.filter(create_by=self.request.user) return qs.filter(create_by=self.request.user)

View File

@ -6,6 +6,8 @@ from apps.utils.filters import MyJsonListFilter
class EquipFilterSet(filters.FilterSet): class EquipFilterSet(filters.FilterSet):
tags = MyJsonListFilter(label='tags/json/list查询') tags = MyJsonListFilter(label='tags/json/list查询')
exclude_cate_name = filters.CharFilter(
field_name='cate__name', exclude=True, label='排除设备分类名称')
class Meta: class Meta:
model = Equipment model = Equipment

View File

@ -1,10 +1,16 @@
from django_filters import rest_framework as filters from django_filters import rest_framework as filters
from apps.inm.models import MaterialBatch, MIO from apps.inm.models import MaterialBatch, MIO
from django.db.models import Q, Subquery, OuterRef, F from django.db.models import Q, Subquery, OuterRef, F
from apps.qm.defect_grades import effective_defect_grade_q
class MaterialBatchFilter(filters.FilterSet): class MaterialBatchFilter(filters.FilterSet):
count_canmio__gt = filters.NumberFilter( count_canmio__gt = filters.NumberFilter(
method='filter_count_canmio__gt', label='可发数量大于') 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: class Meta:
model = MaterialBatch model = MaterialBatch

View File

@ -14,6 +14,7 @@ from django.db.models import F, Sum, DecimalField
from server.settings import get_sysconfig from server.settings import get_sysconfig
from apps.wpmw.models import Wpr from apps.wpmw.models import Wpr
from decimal import Decimal from decimal import Decimal
from apps.qm.defect_grades import DEFECT_GRADE_NAMES, effective_defect_grade
class WareHourseSerializer(CustomModelSerializer): class WareHourseSerializer(CustomModelSerializer):
@ -49,6 +50,8 @@ class MaterialBatchSerializer(CustomModelSerializer):
source='supplier', read_only=True) source='supplier', read_only=True)
material_ = MaterialSerializer(source='material', read_only=True) material_ = MaterialSerializer(source='material', read_only=True)
defect_name = serializers.CharField(source="defect.name", 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='正在出入库数量') count_mioing = serializers.SerializerMethodField(label='正在出入库数量')
class Meta: class Meta:
@ -61,6 +64,12 @@ class MaterialBatchSerializer(CustomModelSerializer):
# 保留 decimal 精度(原 IntegerField 会截断在途量, 导致可发量偏大) # 保留 decimal 精度(原 IntegerField 会截断在途量, 导致可发量偏大)
return instance.count_mioing_anno if hasattr(instance, 'count_mioing_anno') else instance.count_mioing 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): def to_representation(self, instance):
ret = super().to_representation(instance) ret = super().to_representation(instance)
if 'count' in ret: if 'count' in ret:
@ -86,6 +95,15 @@ class MaterialBatchDetailSerializer(CustomModelSerializer):
source='a_mb', read_only=True, many=True) source='a_mb', read_only=True, many=True)
supplier_name = serializers.StringRelatedField( supplier_name = serializers.StringRelatedField(
source='supplier', read_only=True) 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: class Meta:
model = MaterialBatch model = MaterialBatch
@ -542,4 +560,4 @@ class PackSerializer(CustomModelSerializer):
class PackMioSerializer(serializers.Serializer): class PackMioSerializer(serializers.Serializer):
mioitems = serializers.ListField(child=serializers.CharField(), label="明细ID") mioitems = serializers.ListField(child=serializers.CharField(), label="明细ID")
pack_index = serializers.IntegerField(label="包装箱序号") pack_index = serializers.IntegerField(label="包装箱序号")
# pack = serializers.CharField(label="包装箱ID") # pack = serializers.CharField(label="包装箱ID")

View File

@ -4,10 +4,78 @@ from threading import Barrier
from unittest import skipUnless from unittest import skipUnless
from django.db import connection, connections, transaction 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.models import MaterialBatch, WareHouse
from apps.inm.serializers import MaterialBatchSerializer
from apps.mtm.models import Material 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): class MaterialBatchInventoryKeyTests(SimpleTestCase):

View File

@ -60,7 +60,7 @@ class MaterialBatchViewSet(ListModelMixin, CustomGenericViewSet):
queryset = MaterialBatch.objects.filter(count__gt=0) queryset = MaterialBatch.objects.filter(count__gt=0)
serializer_class = MaterialBatchSerializer serializer_class = MaterialBatchSerializer
retrieve_serializer_class = MaterialBatchDetailSerializer retrieve_serializer_class = MaterialBatchDetailSerializer
select_related_fields = ['warehouse', 'material', 'supplier'] select_related_fields = ['warehouse', 'material', 'supplier', 'defect']
filterset_class = MaterialBatchFilter filterset_class = MaterialBatchFilter
search_fields = ['material__name', 'material__number', search_fields = ['material__name', 'material__number',
'material__model', 'material__specification', 'batch'] 'material__model', 'material__specification', 'batch']

42
apps/qm/defect_grades.py Normal file
View File

@ -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

View File

@ -8,19 +8,25 @@ from django.utils.translation import gettext_lazy as _
from django.db import transaction from django.db import transaction
from django.db.models import Sum from django.db.models import Sum
from rest_framework.exceptions import ParseError 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): class Defect(CommonAModel):
"""TN:缺陷项""" """TN:缺陷项"""
DEFECT_OK = 10 DEFECT_OK = GRADE_OK
DEFECT_OK_B = 20 DEFECT_OK_B = GRADE_OK_B
DEFECT_NOTOK = 30 DEFECT_NOTOK = GRADE_NOTOK
cate_list = ["尺寸", "外观", "内质", "性能"] cate_list = ["尺寸", "外观", "内质", "性能"]
name = models.CharField(max_length=50, verbose_name="名称") name = models.CharField(max_length=50, verbose_name="名称")
code = models.CharField(max_length=50, verbose_name="标识", null=True, blank=True) 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)) cate = models.CharField(max_length=50, verbose_name="分类", help_text=str(cate_list))
okcate= models.PositiveSmallIntegerField(verbose_name="不合格分类", okcate= models.PositiveSmallIntegerField(verbose_name="不合格分类",
choices=((DEFECT_OK, "合格"), (DEFECT_OK_B, "合格B类"), (DEFECT_NOTOK, "不合格")), choices=DEFECT_GRADE_CHOICES,
default=DEFECT_NOTOK) default=GRADE_NOTOK)
note = models.TextField('备注', null=True, blank=True) note = models.TextField('备注', null=True, blank=True)
def __str__(self): def __str__(self):

View File

@ -651,6 +651,7 @@ class FileViewSet(BulkCreateModelMixin, RetrieveModelMixin, CustomListModelMixin
class ApkViewSet(MyLoggingMixin, CustomListModelMixin, BulkCreateModelMixin, GenericViewSet): class ApkViewSet(MyLoggingMixin, CustomListModelMixin, BulkCreateModelMixin, GenericViewSet):
perms_map = {'get': '*', 'post': 'apk.upload'} perms_map = {'get': '*', 'post': 'apk.upload'}
serializer_class = ApkSerializer serializer_class = ApkSerializer
filter_backends = []
def get_authenticators(self): def get_authenticators(self):
if self.request.method == 'GET': if self.request.method == 'GET':

View File

@ -69,6 +69,7 @@ class SpeakerViewSet(CustomGenericViewSet):
""" """
perms_map = {} perms_map = {}
serializer_class = serializers.Serializer serializer_class = serializers.Serializer
filter_backends = []
@action(methods=['get'], detail=False, @action(methods=['get'], detail=False,
permission_classes=[IsAuthenticated]) permission_classes=[IsAuthenticated])
@ -125,6 +126,7 @@ class XxTestView(APIView):
class XxCommonViewSet(CreateModelMixin, CustomGenericViewSet): class XxCommonViewSet(CreateModelMixin, CustomGenericViewSet):
perms_map = {'post': '*'} perms_map = {'post': '*'}
serializer_class = RequestCommonSerializer serializer_class = RequestCommonSerializer
filter_backends = []
def create(self, request, *args, **kwargs): def create(self, request, *args, **kwargs):
""" """
@ -258,6 +260,7 @@ class KingCommonViewSet(CreateModelMixin, CustomGenericViewSet):
class DhCommonViewSet(CreateModelMixin, CustomGenericViewSet): class DhCommonViewSet(CreateModelMixin, CustomGenericViewSet):
perms_map = {'post': '*'} perms_map = {'post': '*'}
serializer_class = RequestCommonSerializer serializer_class = RequestCommonSerializer
filter_backends = []
def create(self, request, *args, **kwargs): def create(self, request, *args, **kwargs):
""" """

View File

View File

@ -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}个操作)"
)
)

View File

@ -207,12 +207,28 @@ class CustomRetrieveModelMixin(RetrieveModelMixin):
class CustomListModelMixin(ListModelMixin): class CustomListModelMixin(ListModelMixin):
@swagger_auto_schema(manual_parameters=[ @swagger_auto_schema(
openapi.Parameter(name="query", in_=openapi.IN_QUERY, description="定制返回数据", operation_description=(
type=openapi.TYPE_STRING, required=False), "通用列表接口用于记录或目录浏览以及逐条追溯。跨时间范围的产量、良率、缺陷、"
openapi.Parameter(name="with_children", in_=openapi.IN_QUERY, description="带有children(yes/no/count)", "库存、绩效和趋势等统计聚合,优先查询 BI dataset 目录并执行匹配的数据集。"
type=openapi.TYPE_STRING, required=False), ),
]) 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): def list(self, request, *args, **kwargs):
queryset = self.filter_queryset(self.get_queryset()) queryset = self.filter_queryset(self.get_queryset())

367
apps/utils/swagger.py Normal file
View File

@ -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__

142
apps/utils/test_swagger.py Normal file
View File

@ -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)

View File

@ -154,6 +154,9 @@ class CustomGenericViewSet(MyLoggingMixin, GenericViewSet):
def get_queryset(self): def get_queryset(self):
queryset = super().get_queryset() queryset = super().get_queryset()
queryset = self.get_queryset_custom(queryset) queryset = self.get_queryset_custom(queryset)
# drf-yasg 生成文档时不应读取权限或业务数据。
if getattr(self, 'swagger_fake_view', False):
return queryset
if self.data_filter: if self.data_filter:
user = self.request.user user = self.request.user
if user.is_superuser: if user.is_superuser:
@ -232,4 +235,4 @@ class EuModelViewSet(BulkCreateModelMixin, CustomListModelMixin,
CustomRetrieveModelMixin, BulkDestroyModelMixin, ComplexQueryMixin, CustomGenericViewSet): CustomRetrieveModelMixin, BulkDestroyModelMixin, ComplexQueryMixin, CustomGenericViewSet):
""" """
不支持更新的增强ModelViewSet 不支持更新的增强ModelViewSet
""" """

View File

@ -5,6 +5,7 @@ from apps.mtm.models import Route, Material
from django.db.models import Q, Exists, OuterRef from django.db.models import Q, Exists, OuterRef
from rest_framework.exceptions import ParseError from rest_framework.exceptions import ParseError
from datetime import datetime from datetime import datetime
from apps.qm.defect_grades import effective_defect_grade_q
class SfLogFilter(filters.FilterSet): class SfLogFilter(filters.FilterSet):
class Meta: class Meta:
@ -44,6 +45,10 @@ class WMaterialFilter(filters.FilterSet):
mlog_date_start = filters.DateFilter(label="产出开始", method="filter_mlog_date_start") mlog_date_start = filters.DateFilter(label="产出开始", method="filter_mlog_date_start")
mlog_date_end = filters.DateFilter(label="产出结束", method="filter_mlog_date_end") mlog_date_end = filters.DateFilter(label="产出结束", method="filter_mlog_date_end")
current_merged = filters.BooleanFilter(label="是否本工段新合成的批", method="filter_current_merged") 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): def filter_mlog_date_start(self, queryset, name, value):
mgroupId = self.data.get("mgroup", None) mgroupId = self.data.get("mgroup", None)

View File

@ -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='备注'),
),
]

View File

@ -550,6 +550,7 @@ class MlogUser(BaseModel):
Equipment, verbose_name='生产设备', on_delete=models.CASCADE, null=True, blank=True, related_name='mloguser_equipment') 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) shift = models.ForeignKey(Shift, verbose_name='关联班次', on_delete=models.CASCADE)
handle_date = models.DateField('操作日期') handle_date = models.DateField('操作日期')
note = models.TextField('备注', default='', blank=True)
class Mlogb(BaseModel): class Mlogb(BaseModel):
""" """
@ -876,7 +877,7 @@ class Handoverb(BaseModel):
@property @property
def handoverbw(self): def handoverbw(self):
return Handoverbw.objects.filter(handoverb=self) return self.w_handoverb.all()
class Handoverbw(BaseModel): class Handoverbw(BaseModel):
"""TN: 单个产品交接记录 """TN: 单个产品交接记录

View File

@ -24,12 +24,35 @@ from apps.wpmw.models import Wpr
from apps.qm.serializers import FtestProcessSerializer, FtestProcessListSerializer from apps.qm.serializers import FtestProcessSerializer, FtestProcessListSerializer
import logging import logging
from apps.qm.models import Defect 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 apps.utils.snowflake import idWorker
from decimal import Decimal from decimal import Decimal
from apps.em.models import Equipment from apps.em.models import Equipment
from django.db.models import Q from django.db.models import Q
mylogger = logging.getLogger("log") 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 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 OtherLogSerializer(CustomModelSerializer):
class Meta: class Meta:
model = OtherLog model = OtherLog
@ -199,10 +222,10 @@ class WMaterialSerializer(CustomModelSerializer):
return getattr(NotOkOption, obj.notok_sign, NotOkOption.qt).label if obj.notok_sign else None return getattr(NotOkOption, obj.notok_sign, NotOkOption.qt).label if obj.notok_sign else None
def get_defect_grade(self, obj): 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): 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): def get_count_working(self, obj):
# 列表接口 queryset 已注解(单次聚合); 嵌套等无注解场景回退模型属性 # 列表接口 queryset 已注解(单次聚合); 嵌套等无注解场景回退模型属性
@ -983,10 +1006,18 @@ class MlogbwCreateUpdateSerializer(CustomModelSerializer):
mlogbw = self.save_ftest(mlogbw, ftest_data) mlogbw = self.save_ftest(mlogbw, ftest_data)
return mlogbw return mlogbw
@transaction.atomic
def update(self, instance, validated_data): def update(self, instance, validated_data):
old_number = instance.number
validated_data.pop("mlogb") validated_data.pop("mlogb")
ftest_data = validated_data.pop("ftest", None) ftest_data = validated_data.pop("ftest", None)
mlogbw:Mlogbw = super().update(instance, validated_data) 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: if ftest_data:
mlogbw = self.save_ftest(mlogbw, ftest_data) mlogbw = self.save_ftest(mlogbw, ftest_data)
elif ftest_data is None: elif ftest_data is None:
@ -1256,10 +1287,56 @@ class Handoverbwserializer(CustomModelSerializer):
read_only_fields = EXCLUDE_FIELDS_BASE + ["handoverb", "number"] read_only_fields = EXCLUDE_FIELDS_BASE + ["handoverb", "number"]
extra_kwargs = {'wpr': {'required': True}} 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): 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) 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_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) 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: class Meta:
model = Handoverb model = Handoverb
fields = "__all__" fields = "__all__"
@ -1282,9 +1359,14 @@ class HandoverSerializer(CustomModelSerializer):
recive_user_name = serializers.CharField( recive_user_name = serializers.CharField(
source='recive_user.name', read_only=True) source='recive_user.name', read_only=True)
recive_dept_name = serializers.CharField( 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) send_mgroup_name = serializers.CharField(source='send_mgroup.name', read_only=True)
recive_mgroup_name = serializers.CharField(source='recive_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_ = MaterialSimpleSerializer(source='material', read_only=True)
material_name = serializers.StringRelatedField( material_name = serializers.StringRelatedField(
source='material', read_only=True) source='material', read_only=True)
@ -1292,6 +1374,30 @@ class HandoverSerializer(CustomModelSerializer):
wm_notok_sign = serializers.CharField(source='wm.notok_sign', read_only=True) wm_notok_sign = serializers.CharField(source='wm.notok_sign', read_only=True)
handoverb = HandoverbSerializer(many=True, required=False) handoverb = HandoverbSerializer(many=True, required=False)
ticket_ = TicketSimpleSerializer(source='ticket', read_only=True) 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 = {
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): def validate(self, attrs):
if "mtype" not in attrs: if "mtype" not in attrs:
@ -1410,7 +1516,6 @@ class HandoverSerializer(CustomModelSerializer):
next_mat = None next_mat = None
next_state = None next_state = None
next_defect = None next_defect = None
next_defect_grade = None
if new_wm and attrs["type"] != Handover.H_CHANGE: if new_wm and attrs["type"] != Handover.H_CHANGE:
next_mat = new_wm.material next_mat = new_wm.material
next_state = new_wm.state next_state = new_wm.state
@ -1431,15 +1536,10 @@ class HandoverSerializer(CustomModelSerializer):
if clear_defect and new_wm is not None and new_wm.defect is not None: if clear_defect and new_wm is not None and new_wm.defect is not None:
raise ParseError('清除批次缺陷时目标批次不能带缺陷') raise ParseError('清除批次缺陷时目标批次不能带缺陷')
if clear_defect and tracking == Material.MA_TRACKING_BATCH: 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( raise ParseError(
f'{ind+1}行-批次追踪物料仅同缺陷等级可清除批次缺陷' f'{ind+1}行-批次追踪物料仅合格品和合格B类可清除批次缺陷'
)
if next_defect_grade is None:
next_defect_grade = wm.defect.okcate
elif next_defect_grade != wm.defect.okcate:
raise ParseError(
f'{ind+1}行-批次追踪物料仅同缺陷等级可清除批次缺陷'
) )
if next_mat is None: if next_mat is None:
next_mat = wm.material next_mat = wm.material

View File

@ -1,4 +1,5 @@
import datetime import datetime
from collections import defaultdict
from django.core.cache import cache from django.core.cache import cache
from django.db.models import Sum from django.db.models import Sum
@ -27,6 +28,44 @@ from django.db.models import F
myLogger = logging.getLogger('log') 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): def inherit_zt_batch(source: BatchSt, target: BatchSt):
"""拆批/报工改号时目标批继承来源批的直通统计大批归属(纯继承, 不做判定) """拆批/报工改号时目标批继承来源批的直通统计大批归属(纯继承, 不做判定)

View File

@ -138,13 +138,30 @@ class WMaterialDefectGradeTests(TestCase):
count=1, count=1,
state=WMaterial.WM_OK, 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( normal_notok_data = WMaterialSerializer(
self.normal_with_notok_defect self.normal_with_notok_defect
).data ).data
notok_b_data = WMaterialSerializer(self.notok_with_b_defect).data notok_b_data = WMaterialSerializer(self.notok_with_b_defect).data
no_defect_data = WMaterialSerializer(self.normal_without_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( self.assertEqual(
normal_notok_data["defect_grade"], Defect.DEFECT_NOTOK normal_notok_data["defect_grade"], Defect.DEFECT_NOTOK
@ -154,8 +171,14 @@ class WMaterialDefectGradeTests(TestCase):
notok_b_data["defect_grade"], Defect.DEFECT_OK_B notok_b_data["defect_grade"], Defect.DEFECT_OK_B
) )
self.assertEqual(notok_b_data["defect_grade_name"], "合格B类") self.assertEqual(notok_b_data["defect_grade_name"], "合格B类")
self.assertIsNone(no_defect_data["defect_grade"]) self.assertEqual(no_defect_data["defect_grade"], Defect.DEFECT_OK)
self.assertIsNone(no_defect_data["defect_grade_name"]) 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): def test_filtering_state_and_defect_grade_are_independent(self):
normal_notok = WMaterialFilter( normal_notok = WMaterialFilter(
@ -184,6 +207,28 @@ class WMaterialDefectGradeTests(TestCase):
transform=lambda item: item, 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): class MlogbwViewSetTests(SimpleTestCase):
@patch("apps.wpm.views.MlogViewSet.lock_and_check_can_update") @patch("apps.wpm.views.MlogViewSet.lock_and_check_can_update")
@ -462,7 +507,34 @@ class WMaterialScopeTests(SimpleTestCase):
self.assertTrue(validated["clear_defect"]) self.assertTrue(validated["clear_defect"])
self.assertEqual(validated["count"], 2) 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) material = Material(tracking=Material.MA_TRACKING_BATCH)
defect_a = Defect(id="1", okcate=Defect.DEFECT_NOTOK) defect_a = Defect(id="1", okcate=Defect.DEFECT_NOTOK)
defect_b = Defect(id="2", okcate=Defect.DEFECT_NOTOK) defect_b = Defect(id="2", okcate=Defect.DEFECT_NOTOK)
@ -475,20 +547,21 @@ class WMaterialScopeTests(SimpleTestCase):
state=WMaterial.WM_NOTOK, defect=defect_b, state=WMaterial.WM_NOTOK, defect=defect_b,
) )
validated = HandoverSerializer().validate({ with self.assertRaisesMessage(
"wm": wm_a, ParseError,
"handoverb": [ "批次追踪物料仅合格品和合格B类可清除批次缺陷",
{"wm": wm_a, "count": 1}, ):
{"wm": wm_b, "count": 1}, HandoverSerializer().validate({
], "wm": wm_a,
"new_batch": "N-MERGED", "handoverb": [
"clear_defect": True, {"wm": wm_a, "count": 1},
"type": Handover.H_NORMAL, {"wm": wm_b, "count": 1},
"mtype": Handover.H_MERGE, ],
}) "new_batch": "N-MERGED",
"clear_defect": True,
self.assertTrue(validated["clear_defect"]) "type": Handover.H_NORMAL,
self.assertEqual(validated["count"], 2) "mtype": Handover.H_MERGE,
})
def test_batch_tracking_merge_cannot_clear_mixed_defect_grades(self): def test_batch_tracking_merge_cannot_clear_mixed_defect_grades(self):
material = Material(tracking=Material.MA_TRACKING_BATCH) material = Material(tracking=Material.MA_TRACKING_BATCH)
@ -505,7 +578,7 @@ class WMaterialScopeTests(SimpleTestCase):
with self.assertRaisesMessage( with self.assertRaisesMessage(
ParseError, ParseError,
"批次追踪物料仅同缺陷等级可清除批次缺陷", "批次追踪物料仅合格品和合格B类可清除批次缺陷",
): ):
HandoverSerializer().validate({ HandoverSerializer().validate({
"wm": wm_a, "wm": wm_a,

View File

@ -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,
],
)

View File

@ -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, [])

View File

@ -0,0 +1,87 @@
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, Handoverb, 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)
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)

View File

@ -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")

View File

@ -0,0 +1,193 @@
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.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)
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.values_list.return_value.distinct.return_value.iterator.return_value = []
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)
)
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")

View File

@ -1,12 +1,13 @@
import math import math
import re import re
from string import Formatter
from django.db import transaction from django.db import transaction
from rest_framework.decorators import action from rest_framework.decorators import action
from rest_framework.exceptions import ParseError from rest_framework.exceptions import ParseError
from rest_framework.response import Response from rest_framework.response import Response
from rest_framework.serializers import Serializer 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 django.utils import timezone
from apps.system.models import User from apps.system.models import User
@ -53,13 +54,23 @@ from .serializers import (
MlogUserSerializer, MlogUserSerializer,
BatchLogSerializer, BatchLogSerializer,
MlogQuickSerializer, MlogQuickSerializer,
MlogEquipmentOptionSerializer,
MlogbwStartTestSerializer, MlogbwStartTestSerializer,
HandoverListSerializer, HandoverListSerializer,
BatchChangeSerializer, BatchChangeSerializer,
MlogbOutPatchUpdateSerializer MlogbOutPatchUpdateSerializer
) )
from .services import mlog_submit, handover_submit, mlog_revert, get_batch_dag, handover_revert from .services import (
from apps.wpm.services import mlog_submit_validate, generate_new_batch 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.wf.models import State, Ticket
from apps.wpmw.models import Wpr from apps.wpmw.models import Wpr
from apps.qm.models import Qct, Ftest, TestItem from apps.qm.models import Qct, Ftest, TestItem
@ -73,7 +84,6 @@ from django.db.models import Prefetch
from drf_yasg.utils import swagger_auto_schema from drf_yasg.utils import swagger_auto_schema
from drf_yasg import openapi from drf_yasg import openapi
from django.db import connection from django.db import connection
from django.db.models.functions import Substr, Length
from apps.qm.models import FtestDefect, FtestItem from apps.qm.models import FtestDefect, FtestItem
# Create your views here. # Create your views here.
@ -332,6 +342,92 @@ class MlogViewSet(CustomModelViewSet):
] ]
ordering_fields = ["create_time", "update_time"] 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): def add_info_for_item(self, data):
if data.get("oinfo_json", {}): if data.get("oinfo_json", {}):
czx_dict = dict(TestItem.objects.filter(id__in=data.get("oinfo_json", {}).keys()).values_list("id", "name")) czx_dict = dict(TestItem.objects.filter(id__in=data.get("oinfo_json", {}).keys()).values_list("id", "name"))
@ -353,6 +449,7 @@ class MlogViewSet(CustomModelViewSet):
return super().get_serializer_class() return super().get_serializer_class()
@swagger_auto_schema( @swagger_auto_schema(
operation_summary="查询生产日志明细(逐条追溯)",
manual_parameters=[ manual_parameters=[
openapi.Parameter(name="query", in_=openapi.IN_QUERY, description="定制返回数据", type=openapi.TYPE_STRING, required=False), 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), openapi.Parameter(name="with_children", in_=openapi.IN_QUERY, description="带有children(yes/no/count)", type=openapi.TYPE_STRING, required=False),
@ -591,7 +688,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"] select_related_fields = ["send_user", "send_mgroup", "send_dept", "recive_user", "recive_mgroup", "recive_dept", "wm", "material_changed", "material", "material__process"]
filterset_class = HandoverFilter filterset_class = HandoverFilter
search_fields = ["material__name", "material__number", "material__specification", "batch", "material__model", "b_handover__batch", "new_batch", "wm__batch"] 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): def perform_destroy(self, instance: Handover):
user = self.request.user user = self.request.user
@ -679,7 +789,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) 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]: elif type in [Handover.H_SCRAP]:
m_qs = m_qs.filter(process=None) 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) @action(methods=["post"], detail=False, perms_map={"post": "handover.create"}, serializer_class=GenHandoverWmSerializer)
@transaction.atomic @transaction.atomic
@ -1010,13 +1120,20 @@ class MlogbInViewSet(BulkCreateModelMixin, BulkUpdateModelMixin, BulkDestroyMode
def gen_number_with_rule(cls, rule, material_out: Material, mlog: Mlog, gen_count=1): def gen_number_with_rule(cls, rule, material_out: Material, mlog: Mlog, gen_count=1):
from apps.wpmw.models import Wpr from apps.wpmw.models import Wpr
formatter = Formatter()
rule_parts = list(formatter.parse(rule))
rule_fields = {
field_name
for _, field_name, _, _ in rule_parts
if field_name
}
handle_date = mlog.handle_date handle_date = mlog.handle_date
c_year = handle_date.year c_year = handle_date.year
c_year2 = str(c_year)[-2:] c_year2 = str(c_year)[-2:]
c_month = handle_date.month c_month = handle_date.month
c_day = handle_date.day c_day = handle_date.day
m_model = material_out.model m_model = material_out.model
if 'm_model' in rule: if "m_model" in rule_fields:
if m_model is None: if m_model is None:
raise ParseError("生成编号出错:产品型号不能为空") raise ParseError("生成编号出错:产品型号不能为空")
elif m_model and m_model.islower(): elif m_model and m_model.islower():
@ -1029,29 +1146,64 @@ class MlogbInViewSet(BulkCreateModelMixin, BulkUpdateModelMixin, BulkDestroyMode
if connection.vendor == "postgresql" and connection.in_atomic_block: if connection.vendor == "postgresql" and connection.in_atomic_block:
with connection.cursor() as cursor: with connection.cursor() as cursor:
cursor.execute("SELECT pg_advisory_xact_lock(hashtextextended(%s, 0))", [f"wpr_number_rule:{process.id}"]) cursor.execute("SELECT pg_advisory_xact_lock(hashtextextended(%s, 0))", [f"wpr_number_rule:{process.id}"])
# 按生产日志查询, 流水号归零周期跟随规则中最细的日期占位符 # 只按规则中实际使用的日期占位符筛选历史编号
wpr_filter = { wpr_filter = {
"wpr_mlogbw__mlogb__material_out__isnull": False, "wpr_mlogbw__mlogb__material_out__isnull": False,
"wpr_mlogbw__mlogb__mlog__mgroup__process": process, "wpr_mlogbw__mlogb__mlog__mgroup__process": process,
"wpr_mlogbw__mlogb__mlog__is_fix": False, "wpr_mlogbw__mlogb__mlog__is_fix": False,
"wpr_mlogbw__mlogb__mlog__submit_time__isnull": 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_filter["wpr_mlogbw__mlogb__mlog__handle_date__day"] = c_day
wpr = ( rule_values = {
Wpr.objects.filter(**wpr_filter) "c_year": c_year,
.annotate(last_seq=Substr("number", Length("number") - (cq_w - 1))) "c_year2": c_year2,
.order_by("last_seq") "c_month": c_month,
.last() "c_day": c_day,
) "m_model": m_model,
n_count = 0 }
if wpr: 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: try:
n_count = int(wpr.number[-cq_w:]) field_value = rule_values[field_name]
except Exception as e: if conversion:
raise ParseError(f"获取该类产品最后编号错误: {str(e)}") 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: if n_count + gen_count > 10 ** cq_w - 1:
raise ParseError(f"流水号超出{cq_w}位上限, 请调整编号规则") raise ParseError(f"流水号超出{cq_w}位上限, 请调整编号规则")
try: try:

View File

@ -33,6 +33,19 @@ class Wpr(BaseModel):
data = models.JSONField(verbose_name="数据", default=dict, blank=True) data = models.JSONField(verbose_name="数据", default=dict, blank=True)
pre_info = models.JSONField(verbose_name="预处理信息", default=dict, blank=True, null=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 @classmethod
def change_or_new( def change_or_new(
cls, wpr=None, number=None, mb=None, wm=None, old_mb=None, cls, wpr=None, number=None, mb=None, wm=None, old_mb=None,

View File

@ -63,15 +63,8 @@ class WprViewSet(BulkUpdateModelMixin, CustomListModelMixin, CustomRetrieveModel
vdata = sr.validated_data vdata = sr.validated_data
new_number = vdata["new_number"] new_number = vdata["new_number"]
old_number = vdata["old_number"] old_number = vdata["old_number"]
if Wpr.objects.filter(number=new_number).exists():
raise ParseError("新编号已存在,不可使用")
wpr = Wpr.objects.get(number=old_number) wpr = Wpr.objects.get(number=old_number)
from apps.wpm.models import Mlogbw, Handoverbw wpr.change_number(new_number)
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)
return Response() return Response()
@action(methods=["post"], detail=False, perms_map={"post": "*"}, serializer_class=WprNewSerializer) @action(methods=["post"], detail=False, perms_map={"post": "*"}, serializer_class=WprNewSerializer)

49
docs/mcp.md Normal file
View File

@ -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 <Factory access token>
```
## 配置
生产环境在本机已忽略的 `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。

View File

@ -5,7 +5,12 @@ import sys
def main(): 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: try:
from django.core.management import execute_from_command_line from django.core.management import execute_from_command_line
except ImportError as exc: except ImportError as exc:

1
mcp_server/__init__.py Normal file
View File

@ -0,0 +1 @@
"""Factory MCP v2 integration."""

24
mcp_server/__main__.py Normal file
View File

@ -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()

66
mcp_server/auth.py Normal file
View File

@ -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)

32
mcp_server/context.py Normal file
View File

@ -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}")

89
mcp_server/server.py Normal file
View File

@ -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()

View File

@ -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")

View File

@ -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")

103
mcp_server/test_wprs.py Normal file
View File

@ -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")

164
mcp_server/tests.py Normal file
View File

@ -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,
)

View File

@ -0,0 +1 @@
"""Factory MCP 领域工具,按业务域拆分并在 server 中显式注册。"""

View File

@ -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)

View File

@ -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 响应上限,请缩小查询范围或增加筛选参数"
)

View File

@ -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)

139
mcp_server/tools/wprs.py Normal file
View File

@ -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)

3
pytest.ini Normal file
View File

@ -0,0 +1,3 @@
[pytest]
DJANGO_SETTINGS_MODULE = server.test_settings
python_files = tests.py test_*.py *_tests.py

View File

@ -10,6 +10,11 @@ django-cors-headers==4.9.0
djangorestframework-simplejwt==5.5.1 djangorestframework-simplejwt==5.5.1
django-restql==0.15.2 django-restql==0.15.2
# =======================
# Agent Integration
# =======================
mcp==2.0.0
# ======================= # =======================
# Celery # Celery
# ======================= # =======================

View File

@ -178,6 +178,7 @@ USE_TZ = True
STATIC_URL = '/static/' STATIC_URL = '/static/'
STATIC_ROOT = os.path.join(BASE_DIR, 'dist/static') STATIC_ROOT = os.path.join(BASE_DIR, 'dist/static')
SWAGGER_SCHEMA_PATH = os.path.join(STATIC_ROOT, 'openapi/swagger.json')
# STATICFILES_DIRS = ( # STATICFILES_DIRS = (
# os.path.join(BASE_DIR, 'dist/static'), # os.path.join(BASE_DIR, 'dist/static'),
# ) # )
@ -240,6 +241,20 @@ SIMPLE_JWT = {
'REFRESH_TOKEN_LIFETIME': timedelta(days=60), '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 # 跨域配置/可用nginx处理,无需引入corsheaders
CORS_ORIGIN_ALLOW_ALL = True CORS_ORIGIN_ALLOW_ALL = True
CORS_ALLOW_CREDENTIALS = True CORS_ALLOW_CREDENTIALS = True
@ -267,8 +282,27 @@ CELERYD_SOFT_TIME_LIMIT = 60*10
# swagger配置 # swagger配置
SWAGGER_SETTINGS = { SWAGGER_SETTINGS = {
'DEFAULT_INFO': 'server.swagger.api_info',
'DEFAULT_API_URL': BASE_URL,
'SPEC_URL': 'schema-swagger-json',
'LOGIN_URL': '/django/admin/login/', 'LOGIN_URL': '/django/admin/login/',
'LOGOUT_URL': '/django/admin/logout/', '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 <access token>',
},
'Basic': {
'type': 'basic',
},
},
}
REDOC_SETTINGS = {
'SPEC_URL': 'schema-swagger-json',
} }
# 日志配置 # 日志配置

12
server/swagger.py Normal file
View File

@ -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"),
)

16
server/test_settings.py Normal file
View File

@ -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."
)

View File

@ -17,19 +17,14 @@ from django.conf import settings
from django.conf.urls.static import static from django.conf.urls.static import static
from django.contrib import admin from django.contrib import admin
from django.urls import include, path from django.urls import include, path
from drf_yasg import openapi
from drf_yasg.views import get_schema_view from drf_yasg.views import get_schema_view
from rest_framework.documentation import include_docs_urls from rest_framework.documentation import include_docs_urls
from django.views.generic import TemplateView 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( schema_view = get_schema_view(
openapi.Info( api_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"),
),
public=True, public=True,
permission_classes=[], permission_classes=[],
url=settings.BASE_URL url=settings.BASE_URL
@ -89,8 +84,10 @@ urlpatterns = [
if getattr(settings, 'ENABLE_SWAGGER', True): if getattr(settings, 'ENABLE_SWAGGER', True):
urlpatterns += [ urlpatterns += [
# api文档 # api文档
path('api/swagger.json', swagger_schema_file,
name='schema-swagger-json'),
path('api/swagger/', schema_view.with_ui('swagger', path('api/swagger/', schema_view.with_ui('swagger',
cache_timeout=0), name='schema-swagger-ui'), cache_timeout=0), name='schema-swagger-ui'),
path('api/redoc/', schema_view.with_ui('redoc', path('api/redoc/', schema_view.with_ui('redoc',
cache_timeout=0), name='schema-redoc'), cache_timeout=0), name='schema-redoc'),
] ]