feat(swagger): generate localized static API schema

This commit is contained in:
caoqianming 2026-08-05 16:51:59 +08:00
parent 4bf5f1e585
commit ed952d2d3a
12 changed files with 632 additions and 10 deletions

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

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

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:

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'),
# ) # )
@ -267,8 +268,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"),
)

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,6 +84,8 @@ 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',