fix(develop): restrict unsafe debug endpoints
This commit is contained in:
parent
69b5346031
commit
5fb179eb9e
|
|
@ -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])
|
||||||
|
|
|
||||||
|
|
@ -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))
|
||||||
|
|
|
||||||
|
|
@ -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')
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue