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

View File

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