1
0
forked from erp-dev/erp
Files
erpnew/sse/tests.py

408 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import json
import time
import threading
from queue import Queue
from unittest.mock import patch, MagicMock
from django.test import TestCase
from django.contrib.auth import get_user_model
from django.urls import reverse
from rest_framework import status
from rest_framework.test import APIClient
from rest_framework_simplejwt.tokens import RefreshToken
from basic_info import models as basic_models
from . import services
User = get_user_model()
class SSEAPITestCase(TestCase):
"""SSE API测试用例"""
def setUp(self):
"""设置测试数据"""
# 创建测试商户
self.merchant1 = basic_models.Merchant.objects.create(
name='测试商户1',
type=basic_models.MerchantTypeEnum.STORE,
)
self.merchant2 = basic_models.Merchant.objects.create(
name='测试商户2',
type=basic_models.MerchantTypeEnum.FACTORY,
)
# 创建测试用户
self.user1 = User.objects.create_user(username='user1', password='testpass')
self.user2 = User.objects.create_user(username='user2', password='testpass')
self.user_no_employee = User.objects.create_user(username='no_employee', password='testpass')
# 创建员工并关联商户
self.employee1 = basic_models.Employee.objects.create(
merchant=self.merchant1,
sys_user=self.user1,
name='员工1',
mobile='13800138001',
)
self.employee2 = basic_models.Employee.objects.create(
merchant=self.merchant2,
sys_user=self.user2,
name='员工2',
mobile='13800138002',
)
# 生成JWT token
refresh1 = RefreshToken.for_user(self.user1)
self.token1 = str(refresh1.access_token)
refresh2 = RefreshToken.for_user(self.user2)
self.token2 = str(refresh2.access_token)
refresh_no_employee = RefreshToken.for_user(self.user_no_employee)
self.token_no_employee = str(refresh_no_employee.access_token)
# 设置客户端
self.client = APIClient()
# 清理所有连接
services._connections.clear()
def test_sse_connection_without_token(self):
"""测试没有JWT token的SSE连接"""
response = self.client.get('/sse/')
self.assertEqual(response.status_code, 403)
self.assertEqual(response.content, b'\u8ba4\u8bc1\u5931\u8d25\u6216\u6216\u7528\u6237\u6237')
# 检查连接未被添加到服务中
self.assertEqual(len(services._connections), 0)
def test_sse_connection_with_invalid_token(self):
"""测试无效JWT token的SSE连接"""
response = self.client.get(
'/sse/',
HTTP_AUTHORIZATION='Bearer invalid_token'
)
self.assertEqual(response.status_code, 403)
self.assertEqual(response.content, b'\u8ba4\u8bc1\u5931\u8d25\u6216\u6216\u7528\u6237\u6237')
# 检查连接未被添加到服务中
self.assertEqual(len(services._connections), 0)
def test_sse_connection_user_no_employee(self):
"""测试用户无关联员工的情况"""
response = self.client.get(
'/sse/',
HTTP_AUTHORIZATION=f'Bearer {self.token_no_employee}'
)
self.assertEqual(response.status_code, 403)
self.assertEqual(response.content, b'\u8ba4\u8bc1\u5931\u8d25\u6216\u6216\u7528\u6237\u6237')
# 检查连接未被添加到服务中
self.assertEqual(len(services._connections), 0)
def test_push_test_event_to_merchant(self):
"""测试向特定商户推送测试事件"""
# 手动建立连接
queue1 = Queue()
services.push_connection(self.merchant1.id, queue1)
# 推送事件
response = self.client.post(
'/sse/push/',
HTTP_AUTHORIZATION=f'Bearer {self.token1}'
)
self.assertEqual(response.status_code, 200)
data = response.json()
self.assertEqual(data['status'], 'ok')
self.assertEqual(data['message'], 'Test event broadcasted to your merchant')
self.assertEqual(data['merchant_id'], self.merchant1.id)
self.assertEqual(data['clients'], 1)
# 清理连接
services.remove_connection(self.merchant1.id, queue1)
def test_push_test_event_unauthenticated(self):
"""测试未认证用户推送事件"""
response = self.client.post('/sse/push/')
self.assertEqual(response.status_code, 401)
def test_push_test_event_user_no_employee(self):
"""测试无关联员工用户推送事件"""
response = self.client.post(
'/sse/push/',
HTTP_AUTHORIZATION=f'Bearer {self.token_no_employee}'
)
self.assertEqual(response.status_code, 403)
data = response.json()
self.assertEqual(data['error'], '用户无关联商户')
def test_get_sse_status(self):
"""测试获取SSE状态"""
# 手动建立两个商户的连接
queue1 = Queue()
queue2 = Queue()
services.push_connection(self.merchant1.id, queue1)
services.push_connection(self.merchant2.id, queue2)
# 获取状态
response = self.client.get(
'/sse/status/',
HTTP_AUTHORIZATION=f'Bearer {self.token1}'
)
self.assertEqual(response.status_code, 200)
data = response.json()
self.assertEqual(data['status'], 'running')
self.assertEqual(data['total_clients'], 2) # 所有连接数
self.assertEqual(data['merchant_clients'], 1) # 当前商户的连接数
self.assertEqual(data['merchant_id'], self.merchant1.id)
# 清理连接
services.remove_connection(self.merchant1.id, queue1)
services.remove_connection(self.merchant2.id, queue2)
def test_get_sse_status_unauthenticated(self):
"""测试未认证用户获取状态"""
response = self.client.get('/sse/status/')
self.assertEqual(response.status_code, 401)
def test_shutdown_merchant_connections(self):
"""测试关闭商户连接"""
# 手动建立两个商户的连接
queue1 = Queue()
queue2 = Queue()
services.push_connection(self.merchant1.id, queue1)
services.push_connection(self.merchant2.id, queue2)
# 关闭商户1的连接
response = self.client.post(
'/sse/shutdown/',
HTTP_AUTHORIZATION=f'Bearer {self.token1}'
)
self.assertEqual(response.status_code, 200)
data = response.json()
self.assertEqual(data['status'], 'ok')
self.assertEqual(data['message'], 'Shutdown signal sent to your merchant\'s SSE connections')
self.assertEqual(data['merchant_id'], self.merchant1.id)
self.assertEqual(data['clients'], 1)
# 检查连接状态
merchant1_connections = services.get_merchant_connections(self.merchant1.id)
merchant2_connections = services.get_merchant_connections(self.merchant2.id)
self.assertEqual(len(merchant1_connections), 0)
self.assertEqual(len(merchant2_connections), 1)
# 清理连接
services.remove_connection(self.merchant2.id, queue2)
def test_shutdown_merchant_connections_unauthenticated(self):
"""测试未认证用户关闭连接"""
response = self.client.post('/sse/shutdown/')
self.assertEqual(response.status_code, 401)
def test_merchant_isolation(self):
"""测试商户隔离"""
# 为两个商户分别建立连接
queue1 = Queue()
queue2 = Queue()
queue3 = Queue()
queue4 = Queue()
services.push_connection(self.merchant1.id, queue1)
services.push_connection(self.merchant1.id, queue2)
services.push_connection(self.merchant2.id, queue3)
services.push_connection(self.merchant2.id, queue4)
# 向商户1推送事件
response = self.client.post(
'/sse/push/',
HTTP_AUTHORIZATION=f'Bearer {self.token1}'
)
self.assertEqual(response.status_code, 200)
data = response.json()
self.assertEqual(data['merchant_id'], self.merchant1.id)
self.assertEqual(data['clients'], 2) # 商户1有2个客户端
# 向商户2推送事件
response = self.client.post(
'/sse/push/',
HTTP_AUTHORIZATION=f'Bearer {self.token2}'
)
self.assertEqual(response.status_code, 200)
data = response.json()
self.assertEqual(data['merchant_id'], self.merchant2.id)
self.assertEqual(data['clients'], 2) # 商户2有2个客户端
# 确认商户隔离
merchant1_connections = services.get_merchant_connections(self.merchant1.id)
merchant2_connections = services.get_merchant_connections(self.merchant2.id)
self.assertEqual(len(merchant1_connections), 2)
self.assertEqual(len(merchant2_connections), 2)
# 清理连接
services.remove_connection(self.merchant1.id, queue1)
services.remove_connection(self.merchant1.id, queue2)
services.remove_connection(self.merchant2.id, queue3)
services.remove_connection(self.merchant2.id, queue4)
def test_sse_connection_with_options_request(self):
"""测试OPTIONS预检请求"""
response = self.client.options('/sse/')
# OPTIONS请求应该成功返回CORS头
self.assertEqual(response.status_code, 200)
# 检查CORS头
self.assertIn('Access-Control-Allow-Origin', response)
self.assertIn('Access-Control-Allow-Methods', response)
self.assertIn('Access-Control-Allow-Headers', response)
class SSEServicesTestCase(TestCase):
"""SSE服务层测试用例"""
def setUp(self):
"""设置测试数据"""
self.merchant1_id = 1
self.merchant2_id = 2
# 清理所有连接
services._connections.clear()
def test_push_connection(self):
"""测试添加连接"""
queue1 = Queue()
queue2 = Queue()
# 添加连接
services.push_connection(self.merchant1_id, queue1)
services.push_connection(self.merchant1_id, queue2)
services.push_connection(self.merchant2_id, Queue())
# 检查连接
merchant1_connections = services.get_merchant_connections(self.merchant1_id)
merchant2_connections = services.get_merchant_connections(self.merchant2_id)
self.assertEqual(len(merchant1_connections), 2)
self.assertEqual(len(merchant2_connections), 1)
self.assertIn(queue1, merchant1_connections)
self.assertIn(queue2, merchant1_connections)
def test_remove_connection(self):
"""测试移除连接"""
queue1 = Queue()
queue2 = Queue()
# 添加连接
services.push_connection(self.merchant1_id, queue1)
services.push_connection(self.merchant1_id, queue2)
# 移除一个连接
services.remove_connection(self.merchant1_id, queue1)
# 检查连接
merchant1_connections = services.get_merchant_connections(self.merchant1_id)
self.assertEqual(len(merchant1_connections), 1)
self.assertNotIn(queue1, merchant1_connections)
self.assertIn(queue2, merchant1_connections)
def test_remove_all_merchant_connections(self):
"""测试移除商户所有连接"""
queue1 = Queue()
queue2 = Queue()
# 添加连接
services.push_connection(self.merchant1_id, queue1)
services.push_connection(self.merchant1_id, queue2)
# 移除所有连接
services.remove_connection(self.merchant1_id, queue1)
services.remove_connection(self.merchant1_id, queue2)
# 检查连接
merchant1_connections = services.get_merchant_connections(self.merchant1_id)
self.assertEqual(len(merchant1_connections), 0)
# 商户记录应该被移除
self.assertNotIn(self.merchant1_id, services._connections)
def test_get_all_connections_count(self):
"""测试获取所有连接数"""
# 添加连接
services.push_connection(self.merchant1_id, Queue())
services.push_connection(self.merchant1_id, Queue())
services.push_connection(self.merchant2_id, Queue())
# 检查总连接数
total = services.get_all_connections_count()
self.assertEqual(total, 3)
def test_push_event_to_merchant(self):
"""测试向特定商户推送事件"""
queue1 = Queue()
queue2 = Queue()
queue3 = Queue()
# 添加连接
services.push_connection(self.merchant1_id, queue1)
services.push_connection(self.merchant1_id, queue2)
services.push_connection(self.merchant2_id, queue3)
# 向商户1推送事件
event_data = {'type': 'test', 'message': 'test message'}
services.push_sse_event_to_merchant(self.merchant1_id, event_data)
# 检查消息
self.assertEqual(queue1.qsize(), 1)
self.assertEqual(queue2.qsize(), 1)
self.assertEqual(queue3.qsize(), 0) # 商户2不应该收到消息
# 检查消息内容
self.assertEqual(queue1.get_nowait(), event_data)
self.assertEqual(queue2.get_nowait(), event_data)
def test_push_event_to_nonexistent_merchant(self):
"""测试向不存在的商户推送事件"""
# 向不存在的商户推送事件
event_data = {'type': 'test', 'message': 'test message'}
services.push_sse_event_to_merchant(999, event_data)
# 不应该抛出异常,也不会有连接受到影响
self.assertEqual(len(services._connections), 0)
def test_push_simple_message_to_merchant(self):
"""测试向特定商户推送简单消息"""
queue1 = Queue()
# 添加连接
services.push_connection(self.merchant1_id, queue1)
# 推送简单消息
services.push_simple_message_with_object_id_to_merchant(
self.merchant1_id,
'order_paid',
'订单已支付',
12345
)
# 检查消息
self.assertEqual(queue1.qsize(), 1)
# 检查消息内容
message = queue1.get_nowait()
self.assertEqual(message['mode'], 'simple_message')
self.assertEqual(message['type'], 'order_paid')
self.assertEqual(message['message'], '订单已支付')
self.assertEqual(message['object_id'], 12345)
self.assertEqual(message['merchant_id'], self.merchant1_id)