From 5fb179eb9e678877be785388ef8ebc8c05640b19 Mon Sep 17 00:00:00 2001 From: caoqianming Date: Tue, 4 Aug 2026 16:33:27 +0800 Subject: [PATCH] fix(develop): restrict unsafe debug endpoints --- apps/develop/tests.py | 33 +++++++++++++++++++++++++++++++-- apps/develop/urls.py | 11 ++++++++--- apps/develop/views.py | 27 +++++++++++++-------------- 3 files changed, 52 insertions(+), 19 deletions(-) diff --git a/apps/develop/tests.py b/apps/develop/tests.py index 7ce503c2..9a7f0fe3 100755 --- a/apps/develop/tests.py +++ b/apps/develop/tests.py @@ -1,3 +1,32 @@ -from django.test import TestCase +from django.test import SimpleTestCase +from rest_framework.permissions import IsAdminUser +from rest_framework.test import APIRequestFactory -# Create your tests here. +from apps.develop.views import ServerTime, TestViewSet + + +class DevelopApiPermissionTests(SimpleTestCase): + def setUp(self): + self.factory = APIRequestFactory() + + def test_test_endpoint_rejects_anonymous_requests(self): + request = self.factory.post( + '/api/develop/test/send_sms/', + {}, + format='json', + ) + + response = TestViewSet.as_view({'post': 'send_sms'})(request) + + self.assertIn(response.status_code, (401, 403)) + + def test_server_time_requires_admin(self): + request = self.factory.get('/api/develop/server_time/') + + response = ServerTime.as_view()(request) + + self.assertIn(response.status_code, (401, 403)) + + def test_develop_views_use_admin_permission(self): + self.assertEqual(TestViewSet.permission_classes, [IsAdminUser]) + self.assertEqual(ServerTime.permission_classes, [IsAdminUser]) diff --git a/apps/develop/urls.py b/apps/develop/urls.py index bd894ba8..1986082d 100755 --- a/apps/develop/urls.py +++ b/apps/develop/urls.py @@ -1,4 +1,5 @@ -from django.urls import path, include +from django.conf import settings +from django.urls import include, path from apps.develop.views import (BackupDatabase, BackupMedia, ReloadClientGit, ReloadServerGit, ReloadServerOnly, TestViewSet, CorrectViewSet, testScanHtml, ServerTime) from rest_framework.routers import DefaultRouter @@ -6,9 +7,11 @@ from rest_framework.routers import DefaultRouter API_BASE_URL = 'api/develop/' HTML_BASE_URL = 'dhtml/develop/' router = DefaultRouter() -router.register('test', TestViewSet, basename='api_test') router.register('correct', CorrectViewSet, basename='correct') +if settings.DEBUG: + router.register('test', TestViewSet, basename='api_test') + urlpatterns = [ path(API_BASE_URL + 'reload_server_git/', ReloadServerGit.as_view()), # path(API_BASE_URL + 'reload_web_git/', ReloadClientGit.as_view()), @@ -17,5 +20,7 @@ urlpatterns = [ path(API_BASE_URL + 'backup_media/', BackupMedia.as_view()), path(API_BASE_URL + 'server_time/', ServerTime.as_view()), path(API_BASE_URL, include(router.urls)), - path(HTML_BASE_URL + "testscan/", testScanHtml) ] + +if settings.DEBUG: + urlpatterns.append(path(HTML_BASE_URL + "testscan/", testScanHtml)) diff --git a/apps/develop/views.py b/apps/develop/views.py index 03e493da..0c4c2fbe 100755 --- a/apps/develop/views.py +++ b/apps/develop/views.py @@ -2,7 +2,7 @@ from rest_framework.views import APIView from rest_framework.exceptions import ParseError -from rest_framework.permissions import IsAdminUser, AllowAny +from rest_framework.permissions import IsAdminUser from rest_framework.response import Response from rest_framework.serializers import Serializer from rest_framework.decorators import action @@ -40,11 +40,7 @@ from datetime import datetime # Create your views here. class ServerTime(APIView): - - def get_permissions(self): - if self.request.method == 'GET': - return [AllowAny()] - return [IsAdminUser()] + permission_classes = [IsAdminUser] @swagger_auto_schema(responses={200: ServerTimeSerializer}) def get(self, request): @@ -62,9 +58,13 @@ class ServerTime(APIView): 修改服务器时间 """ - command = f'date -s "{request.data["server_time"]}"' + serializer = ServerTimeSerializer(data=request.data) + serializer.is_valid(raise_exception=True) + server_time = serializer.validated_data['server_time'].strftime( + "%Y-%m-%d %H:%M:%S" + ) completed = subprocess.run( - ["sudo", "-S", "sh", "-c", command], # 添加 -S 参数 + ["sudo", "-S", "date", "-s", server_time], input=SD_PWD + "\n", # 注意要在密码后加换行符 capture_output=True, text=True @@ -269,10 +269,9 @@ class CorrectViewSet(CustomGenericViewSet): class TestViewSet(CustomGenericViewSet): perms_map = {} - authentication_classes = () - permission_classes = () + permission_classes = [IsAdminUser] - @action(methods=['post'], detail=False, serializer_class=SendSmsSerializer, authentication_classes=()) + @action(methods=['post'], detail=False, serializer_class=SendSmsSerializer) def send_sms(self, request, pk=None): """发送短信测试 @@ -565,7 +564,7 @@ class TestViewSet(CustomGenericViewSet): # correct_card_time() # return Response() - @action(methods=['post'], detail=False, serializer_class=Serializer, permission_classes=[]) + @action(methods=['post'], detail=False, serializer_class=Serializer) @transaction.atomic def correct_data(self, request, pk=None): """修正数据 @@ -674,7 +673,7 @@ class TestViewSet(CustomGenericViewSet): Ticket.objects.get_queryset(all=True).delete() return Response() - @action(methods=['post'], detail=False, serializer_class=Serializer, permission_classes=[]) + @action(methods=['post'], detail=False, serializer_class=Serializer) def test_cal(self, request, pk=None): from apps.wpm.tasks import cal_exp_duration_sec cal_exp_duration_sec('3397169058570170368') @@ -710,4 +709,4 @@ html_str = """ """ def testScanHtml(request): - return HttpResponse(html_str) \ No newline at end of file + return HttpResponse(html_str)