fix(develop): restrict unsafe debug endpoints

This commit is contained in:
caoqianming 2026-08-04 16:33:27 +08:00
parent 69b5346031
commit 5fb179eb9e
3 changed files with 52 additions and 19 deletions

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