diff --git a/tests/unit/test_bind_contact_use_case.py b/tests/unit/test_bind_contact_use_case.py new file mode 100755 index 000000000..3e91021ac --- /dev/null +++ b/tests/unit/test_bind_contact_use_case.py @@ -0,0 +1,484 @@ +"""绑定联系方式 UseCase 单元测试.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +from packages.application.auth.bind_contact_use_case import ( + BindContactRequest, + BindContactUseCase, + SendVerificationCodeRequest, + SendVerificationCodeUseCase, +) +from packages.domain.entities import User + + +@pytest.fixture +def mock_user_repo(): + return MagicMock() + + +@pytest.fixture +def mock_verification_service(): + svc = MagicMock() + svc.verify.return_value = (True, None) + return svc + + +@pytest.fixture +def sample_user(): + user = User( + id="user_001", + email="", + display_name="测试用户", + phone_verified=False, + email_verified=False, + ) + user.phone = None + return user + + +class TestBindContactRequest: + """BindContactRequest 测试""" + + def test_phone_strips_plus86(self): + """手机号 +86 前缀会被去掉""" + req = BindContactRequest( + user_id="u1", phone="+8613800000001", phone_code="1234" + ) + assert req.phone == "13800000001" + + def test_email_lowercased(self): + """邮箱会被转小写""" + req = BindContactRequest( + user_id="u1", email="Test@Example.COM", email_code="1234" + ) + assert req.email == "test@example.com" + + def test_code_stripped(self): + """验证码会被 strip""" + req = BindContactRequest( + user_id="u1", phone="13800000001", phone_code=" 1234 " + ) + assert req.phone_code == "1234" + + def test_empty_fields(self): + """空字段处理""" + req = BindContactRequest(user_id="u1") + assert req.phone == "" + assert req.email == "" + assert req.phone_code == "" + assert req.email_code == "" + + +class TestBindContactUseCase: + """BindContactUseCase 测试""" + + def test_bind_phone_success(self, mock_user_repo, mock_verification_service, sample_user): + """绑定手机号成功""" + mock_user_repo.find_by_id.return_value = sample_user + mock_user_repo.find_by_phone.return_value = None + mock_user_repo.save.return_value = sample_user + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + phone="13800000001", + phone_code="123456", + ) + response, error = use_case.execute(request) + + assert error is None + assert response is not None + assert response.user.phone == "13800000001" + assert response.user.phone_verified is True + mock_user_repo.save.assert_called_once() + + def test_bind_email_success(self, mock_user_repo, mock_verification_service, sample_user): + """绑定邮箱成功""" + mock_user_repo.find_by_id.return_value = sample_user + mock_user_repo.find_by_email.return_value = None + mock_user_repo.save.return_value = sample_user + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + email="test@example.com", + email_code="123456", + ) + response, error = use_case.execute(request) + + assert error is None + assert response is not None + assert response.user.email == "test@example.com" + assert response.user.email_verified is True + + def test_bind_phone_and_email(self, mock_user_repo, mock_verification_service, sample_user): + """同时绑定手机和邮箱""" + mock_user_repo.find_by_id.return_value = sample_user + mock_user_repo.find_by_phone.return_value = None + mock_user_repo.find_by_email.return_value = None + mock_user_repo.save.return_value = sample_user + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + phone="13800000001", + phone_code="123456", + email="test@example.com", + email_code="123456", + ) + response, error = use_case.execute(request) + + assert error is None + assert response.user.phone == "13800000001" + assert response.user.phone_verified is True + assert response.user.email == "test@example.com" + assert response.user.email_verified is True + # 两个都绑定完成,binding_completed_at 应该被设置 + assert response.user.binding_completed_at is not None + + def test_no_contact_info_returns_error(self, mock_user_repo, mock_verification_service): + """既没填手机也没填邮箱返回错误""" + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest(user_id="user_001") + + response, error = use_case.execute(request) + + assert response is None + assert "至少填写" in error + mock_user_repo.find_by_id.assert_not_called() + + def test_user_not_found(self, mock_user_repo, mock_verification_service): + """用户不存在返回错误""" + mock_user_repo.find_by_id.return_value = None + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="nonexistent", + phone="13800000001", + phone_code="123456", + ) + + response, error = use_case.execute(request) + + assert response is None + assert "用户不存在" in error + + def test_phone_already_bound_by_other(self, mock_user_repo, mock_verification_service, sample_user): + """手机号已被其他账号绑定""" + other_user = MagicMock() + other_user.id = "user_other" + mock_user_repo.find_by_id.return_value = sample_user + mock_user_repo.find_by_phone.return_value = other_user + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + phone="13800000001", + phone_code="123456", + ) + + response, error = use_case.execute(request) + + assert response is None + assert "已被其他账号绑定" in error + mock_user_repo.save.assert_not_called() + + def test_phone_bound_by_self_ok(self, mock_user_repo, mock_verification_service, sample_user): + """手机号已被自己绑定,允许""" + sample_user.phone = "13800000001" + mock_user_repo.find_by_id.return_value = sample_user + mock_user_repo.find_by_phone.return_value = sample_user + mock_user_repo.save.return_value = sample_user + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + phone="13800000001", + phone_code="123456", + ) + + response, error = use_case.execute(request) + assert error is None + assert response is not None + + def test_wrong_phone_code(self, mock_user_repo, mock_verification_service, sample_user): + """手机验证码错误""" + mock_user_repo.find_by_id.return_value = sample_user + mock_user_repo.find_by_phone.return_value = None + mock_verification_service.verify.return_value = (False, "验证码过期") + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + phone="13800000001", + phone_code="000000", + ) + + response, error = use_case.execute(request) + + assert response is None + assert "手机验证码错误" in error + mock_user_repo.save.assert_not_called() + + def test_missing_phone_code(self, mock_user_repo, mock_verification_service, sample_user): + """缺少手机验证码""" + mock_user_repo.find_by_id.return_value = sample_user + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + phone="13800000001", + phone_code="", + ) + + response, error = use_case.execute(request) + + assert response is None + assert "请输入手机验证码" in error + + def test_invalid_phone_format(self, mock_user_repo, mock_verification_service, sample_user): + """手机号格式不正确""" + mock_user_repo.find_by_id.return_value = sample_user + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + phone="123", # 太短 + phone_code="123456", + ) + + response, error = use_case.execute(request) + + assert response is None + assert error is not None + + def test_email_already_bound_by_other(self, mock_user_repo, mock_verification_service, sample_user): + """邮箱已被其他账号绑定""" + other_user = MagicMock() + other_user.id = "user_other" + mock_user_repo.find_by_id.return_value = sample_user + mock_user_repo.find_by_email.return_value = other_user + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + email="test@example.com", + email_code="123456", + ) + + response, error = use_case.execute(request) + + assert response is None + assert "已被其他账号绑定" in error + + def test_missing_email_code(self, mock_user_repo, mock_verification_service, sample_user): + """缺少邮箱验证码""" + mock_user_repo.find_by_id.return_value = sample_user + mock_user_repo.find_by_email.return_value = None + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + email="test@example.com", + email_code="", + ) + + response, error = use_case.execute(request) + + assert response is None + assert "请输入邮箱验证码" in error + + def test_invalid_email_format(self, mock_user_repo, mock_verification_service, sample_user): + """邮箱格式不正确""" + mock_user_repo.find_by_id.return_value = sample_user + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + email="not_an_email", + email_code="123456", + ) + + response, error = use_case.execute(request) + + assert response is None + assert error is not None + + def test_response_to_dict(self, mock_user_repo, mock_verification_service, sample_user): + """BindContactResponse.to_dict 返回正确格式""" + mock_user_repo.find_by_id.return_value = sample_user + mock_user_repo.find_by_phone.return_value = None + mock_user_repo.find_by_email.return_value = None + mock_user_repo.save.return_value = sample_user + + use_case = BindContactUseCase(mock_user_repo, mock_verification_service) + request = BindContactRequest( + user_id="user_001", + phone="13800000001", + phone_code="123456", + email="test@example.com", + email_code="123456", + ) + response, _ = use_case.execute(request) + data = response.to_dict() + + assert "user" in data + assert data["user"]["id"] == "user_001" + assert "email" in data["user"] + assert "phone" in data["user"] + assert "phone_verified" in data["user"] + assert "display_name" in data["user"] + assert "binding_complete" in data["user"] + + +class TestSendVerificationCodeRequest: + """SendVerificationCodeRequest 测试""" + + def test_value_stripped(self): + """value 会被 strip""" + req = SendVerificationCodeRequest(target="phone", value=" 13800000001 ", purpose="bind") + assert req.value == "13800000001" + + +class TestSendVerificationCodeUseCase: + """SendVerificationCodeUseCase 测试""" + + def test_send_phone_code_success(self, mock_verification_service): + """发送手机验证码成功""" + from datetime import datetime, timedelta, timezone + + code_obj = MagicMock() + code_obj.code = "123456" + code_obj.created_at = datetime.now(timezone.utc) + code_obj.expires_at = datetime.now(timezone.utc) + timedelta(minutes=5) + mock_verification_service.generate.return_value = (code_obj, None) + + mock_sms = MagicMock() + use_case = SendVerificationCodeUseCase( + mock_verification_service, + sms_service=mock_sms, + ) + request = SendVerificationCodeRequest( + target="phone", + value="13800000001", + purpose="bind", + ) + response, error = use_case.execute(request) + + assert error is None + assert response is not None + assert response.expires_in > 0 + assert response.resend_after == 60 + mock_sms.send_verification_code.assert_called_once() + + def test_send_email_code_success(self, mock_verification_service): + """发送邮箱验证码成功""" + from datetime import datetime, timedelta, timezone + + code_obj = MagicMock() + code_obj.code = "654321" + code_obj.created_at = datetime.now(timezone.utc) + code_obj.expires_at = datetime.now(timezone.utc) + timedelta(minutes=5) + mock_verification_service.generate.return_value = (code_obj, None) + + mock_email = MagicMock() + use_case = SendVerificationCodeUseCase( + mock_verification_service, + email_service=mock_email, + ) + request = SendVerificationCodeRequest( + target="email", + value="test@example.com", + purpose="bind", + ) + response, error = use_case.execute(request) + + assert error is None + assert response is not None + mock_email.send_email.assert_called_once() + + def test_invalid_target_returns_error(self, mock_verification_service): + """不支持的目标类型返回错误""" + use_case = SendVerificationCodeUseCase(mock_verification_service) + request = SendVerificationCodeRequest( + target="wechat", + value="some_value", + purpose="bind", + ) + + response, error = use_case.execute(request) + + assert response is None + assert "不支持的目标类型" in error + + def test_invalid_phone_format(self, mock_verification_service): + """手机号格式错误返回错误""" + use_case = SendVerificationCodeUseCase(mock_verification_service) + request = SendVerificationCodeRequest( + target="phone", + value="123", + purpose="bind", + ) + + response, error = use_case.execute(request) + + assert response is None + assert error is not None + mock_verification_service.generate.assert_not_called() + + def test_invalid_email_format(self, mock_verification_service): + """邮箱格式错误返回错误""" + use_case = SendVerificationCodeUseCase(mock_verification_service) + request = SendVerificationCodeRequest( + target="email", + value="not_email", + purpose="bind", + ) + + response, error = use_case.execute(request) + + assert response is None + assert error is not None + mock_verification_service.generate.assert_not_called() + + def test_generate_failure_returns_error(self, mock_verification_service): + """生成验证码失败返回错误""" + mock_verification_service.generate.return_value = (None, "发送太频繁") + + use_case = SendVerificationCodeUseCase(mock_verification_service) + request = SendVerificationCodeRequest( + target="phone", + value="13800000001", + purpose="bind", + ) + + response, error = use_case.execute(request) + + assert response is None + assert "发送太频繁" in error + + def test_response_to_dict(self, mock_verification_service): + """SendVerificationCodeResponse.to_dict 格式正确""" + from datetime import datetime, timedelta, timezone + + code_obj = MagicMock() + code_obj.code = "123456" + code_obj.created_at = datetime.now(timezone.utc) + code_obj.expires_at = datetime.now(timezone.utc) + timedelta(seconds=300) + mock_verification_service.generate.return_value = (code_obj, None) + + use_case = SendVerificationCodeUseCase(mock_verification_service) + request = SendVerificationCodeRequest( + target="phone", + value="13800000001", + purpose="bind", + ) + response, _ = use_case.execute(request) + data = response.to_dict() + + assert "expires_in" in data + assert "resend_after" in data diff --git a/tests/unit/test_pagination.py b/tests/unit/test_pagination.py index 37ee01da8..9d9dcc7c1 100755 --- a/tests/unit/test_pagination.py +++ b/tests/unit/test_pagination.py @@ -1,12 +1,6 @@ -""" -pagination 通用分页器单元测试 +"""通用分页器单元测试.""" -覆盖: -- PaginationParams: 默认值/边界/校验/offset/limit -- PaginationMeta: from_params 各种边界场景 -- PaginatedResponse: create 工厂方法 -- paginate: 内存分页函数 -""" +from __future__ import annotations import pytest from pydantic import ValidationError @@ -18,320 +12,233 @@ from packages.application.common.pagination import ( paginate, ) -# ============================================================ -# PaginationParams -# ============================================================ +class TestPaginationParams: + """PaginationParams 测试""" -class TestPaginationParamsDefaults: - """默认值测试""" - - def test_default_page_is_1(self): + def test_default_values(self): + """默认值正确""" params = PaginationParams() assert params.page == 1 - - def test_default_page_size_is_20(self): - params = PaginationParams() assert params.page_size == 20 - def test_default_offset_is_0(self): - params = PaginationParams() + def test_offset_first_page(self): + """第一页 offset 为 0""" + params = PaginationParams(page=1, page_size=20) assert params.offset == 0 - def test_default_limit_is_20(self): - params = PaginationParams() - assert params.limit == 20 + def test_offset_second_page(self): + """第二页 offset 计算正确""" + params = PaginationParams(page=2, page_size=20) + assert params.offset == 20 + def test_offset_custom_page_size(self): + """自定义 page_size 的 offset""" + params = PaginationParams(page=3, page_size=10) + assert params.offset == 20 -class TestPaginationParamsValidation: - """参数校验""" + def test_limit_equals_page_size(self): + """limit 等于 page_size""" + params = PaginationParams(page_size=50) + assert params.limit == 50 - @pytest.mark.parametrize("page", [1, 2, 100, 9999]) - def test_valid_page_values(self, page): - params = PaginationParams(page=page) - assert params.page == page - - def test_page_zero_raises(self): + def test_page_must_be_at_least_1(self): + """page 不能小于 1""" with pytest.raises(ValidationError): PaginationParams(page=0) def test_page_negative_raises(self): + """page 不能为负数""" with pytest.raises(ValidationError): PaginationParams(page=-1) - @pytest.mark.parametrize("page_size", [1, 20, 50, 100]) - def test_valid_page_size_values(self, page_size): - params = PaginationParams(page_size=page_size) - assert params.page_size == page_size - - def test_page_size_zero_raises(self): + def test_page_size_must_be_at_least_1(self): + """page_size 不能小于 1""" with pytest.raises(ValidationError): PaginationParams(page_size=0) - def test_page_size_negative_raises(self): - with pytest.raises(ValidationError): - PaginationParams(page_size=-5) - - def test_page_size_over_100_raises(self): + def test_page_size_max_100(self): + """page_size 最大 100""" with pytest.raises(ValidationError): PaginationParams(page_size=101) - def test_invalid_page_type_raises(self): - with pytest.raises(ValidationError): - PaginationParams(page="abc") - - def test_invalid_page_size_type_raises(self): - with pytest.raises(ValidationError): - PaginationParams(page_size="abc") - - -class TestPaginationParamsOffset: - """offset 属性计算""" - - def test_page_1_offset_0(self): - params = PaginationParams(page=1, page_size=20) - assert params.offset == 0 - - def test_page_2_offset_page_size(self): - params = PaginationParams(page=2, page_size=20) - assert params.offset == 20 - - def test_page_3_offset_2x_page_size(self): - params = PaginationParams(page=3, page_size=20) - assert params.offset == 40 - - def test_page_5_page_size_10_offset_40(self): - params = PaginationParams(page=5, page_size=10) - assert params.offset == 40 - - def test_page_1_page_size_100_offset_0(self): - params = PaginationParams(page=1, page_size=100) - assert params.offset == 0 - - -class TestPaginationParamsLimit: - """limit 属性""" - - def test_limit_equals_page_size(self): - params = PaginationParams(page_size=20) - assert params.limit == 20 - - def test_limit_1(self): - params = PaginationParams(page_size=1) - assert params.limit == 1 - - def test_limit_100(self): + def test_page_size_100_is_valid(self): + """page_size=100 是合法的""" params = PaginationParams(page_size=100) - assert params.limit == 100 + assert params.page_size == 100 -# ============================================================ -# PaginationMeta.from_params -# ============================================================ +class TestPaginationMeta: + """PaginationMeta 测试""" + def test_from_params_first_page(self): + """第一页元数据""" + params = PaginationParams(page=1, page_size=10) + meta = PaginationMeta.from_params(params, total=25) -class TestPaginationMetaFromParams: - """from_params 工厂方法""" + assert meta.page == 1 + assert meta.page_size == 10 + assert meta.total == 25 + assert meta.total_pages == 3 + assert meta.has_next is True + assert meta.has_prev is False - def test_empty_total_zero(self): + def test_from_params_last_page(self): + """最后一页元数据""" + params = PaginationParams(page=3, page_size=10) + meta = PaginationMeta.from_params(params, total=25) + + assert meta.page == 3 + assert meta.total_pages == 3 + assert meta.has_next is False + assert meta.has_prev is True + + def test_from_params_middle_page(self): + """中间页元数据""" + params = PaginationParams(page=2, page_size=10) + meta = PaginationMeta.from_params(params, total=50) + + assert meta.page == 2 + assert meta.total_pages == 5 + assert meta.has_next is True + assert meta.has_prev is True + + def test_from_params_zero_total(self): + """总数为 0 时""" params = PaginationParams(page=1, page_size=20) meta = PaginationMeta.from_params(params, total=0) + assert meta.total == 0 assert meta.total_pages == 0 assert meta.has_next is False assert meta.has_prev is False - def test_exactly_one_page(self): - params = PaginationParams(page=1, page_size=20) - meta = PaginationMeta.from_params(params, total=20) + def test_from_params_exact_multiple(self): + """总数刚好是 page_size 的整数倍""" + params = PaginationParams(page=1, page_size=10) + meta = PaginationMeta.from_params(params, total=30) + + assert meta.total_pages == 3 + + def test_from_params_single_page(self): + """单页即可放下所有数据""" + params = PaginationParams(page=1, page_size=100) + meta = PaginationMeta.from_params(params, total=50) + assert meta.total_pages == 1 assert meta.has_next is False assert meta.has_prev is False - def test_less_than_one_page(self): - params = PaginationParams(page=1, page_size=20) - meta = PaginationMeta.from_params(params, total=15) - assert meta.total_pages == 1 - assert meta.has_next is False - assert meta.has_prev is False - def test_multiple_pages_first_page(self): - params = PaginationParams(page=1, page_size=20) - meta = PaginationMeta.from_params(params, total=50) - assert meta.total_pages == 3 - assert meta.has_next is True - assert meta.has_prev is False +class TestPaginatedResponse: + """PaginatedResponse 测试""" - def test_multiple_pages_middle_page(self): - params = PaginationParams(page=2, page_size=20) - meta = PaginationMeta.from_params(params, total=50) - assert meta.total_pages == 3 - assert meta.has_next is True - assert meta.has_prev is True - - def test_multiple_pages_last_page(self): - params = PaginationParams(page=3, page_size=20) - meta = PaginationMeta.from_params(params, total=50) - assert meta.total_pages == 3 - assert meta.has_next is False - assert meta.has_prev is True - - def test_exact_division(self): - params = PaginationParams(page=2, page_size=20) - meta = PaginationMeta.from_params(params, total=40) - assert meta.total_pages == 2 - assert meta.has_next is False - assert meta.has_prev is True - - def test_non_exact_division_ceil(self): - params = PaginationParams(page=1, page_size=20) - meta = PaginationMeta.from_params(params, total=41) - assert meta.total_pages == 3 - - def test_total_1_page_size_20(self): - params = PaginationParams(page=1, page_size=20) - meta = PaginationMeta.from_params(params, total=1) - assert meta.total_pages == 1 - assert meta.has_next is False - assert meta.has_prev is False - - def test_page_beyond_total_pages(self): - params = PaginationParams(page=10, page_size=20) - meta = PaginationMeta.from_params(params, total=50) - assert meta.total_pages == 3 - assert meta.has_next is False - assert meta.has_prev is True - - def test_preserves_params_values(self): - params = PaginationParams(page=3, page_size=15) - meta = PaginationMeta.from_params(params, total=100) - assert meta.page == 3 - assert meta.page_size == 15 - assert meta.total == 100 - - -# ============================================================ -# PaginatedResponse.create -# ============================================================ - - -class TestPaginatedResponseCreate: - """create 工厂方法""" - - def test_create_with_data(self): - params = PaginationParams(page=1, page_size=20) + def test_create_success(self): + """创建分页响应""" + params = PaginationParams(page=1, page_size=10) data = [1, 2, 3] - response = PaginatedResponse.create(data, params, total=100) - assert response.data == data - assert response.pagination.total == 100 - assert response.pagination.page == 1 - assert response.pagination.page_size == 20 - def test_create_with_empty_data(self): + response = PaginatedResponse.create(data, params, total=25) + + assert response.data == [1, 2, 3] + assert response.pagination.page == 1 + assert response.pagination.total == 25 + assert response.pagination.total_pages == 3 + + def test_create_empty_data(self): + """空数据分页响应""" params = PaginationParams(page=1, page_size=20) response = PaginatedResponse.create([], params, total=0) + assert response.data == [] assert response.pagination.total == 0 assert response.pagination.total_pages == 0 - def test_create_preserves_list_type(self): - params = PaginationParams(page=1, page_size=20) - data = ["a", "b", "c"] - response = PaginatedResponse.create(data, params, total=10) - assert response.data == ["a", "b", "c"] - assert len(response.data) == 3 - - -# ============================================================ -# paginate 函数 -# ============================================================ - class TestPaginateFunction: - """内存分页函数""" - - def test_empty_list(self): - params = PaginationParams(page=1, page_size=20) - result = paginate([], params) - assert result.data == [] - assert result.pagination.total == 0 - assert result.pagination.total_pages == 0 + """paginate 函数测试(内存分页)""" def test_first_page(self): - items = list(range(50)) - params = PaginationParams(page=1, page_size=20) + """第一页分页""" + items = list(range(30)) + params = PaginationParams(page=1, page_size=10) + result = paginate(items, params) - assert result.data == list(range(20)) - assert result.pagination.total == 50 + + assert result.data == list(range(10)) + assert result.pagination.total == 30 assert result.pagination.total_pages == 3 assert result.pagination.has_next is True assert result.pagination.has_prev is False - def test_middle_page(self): - items = list(range(50)) - params = PaginationParams(page=2, page_size=20) + def test_second_page(self): + """第二页分页""" + items = list(range(30)) + params = PaginationParams(page=2, page_size=10) + result = paginate(items, params) - assert result.data == list(range(20, 40)) - assert result.pagination.has_next is True - assert result.pagination.has_prev is True + + assert result.data == list(range(10, 20)) + assert result.pagination.page == 2 def test_last_page(self): - items = list(range(50)) - params = PaginationParams(page=3, page_size=20) + """最后一页分页""" + items = list(range(25)) + params = PaginationParams(page=3, page_size=10) + result = paginate(items, params) - assert result.data == list(range(40, 50)) - assert len(result.data) == 10 + + assert result.data == list(range(20, 25)) + assert len(result.data) == 5 assert result.pagination.has_next is False - assert result.pagination.has_prev is True + + def test_empty_list(self): + """空列表分页""" + params = PaginationParams(page=1, page_size=20) + result = paginate([], params) + + assert result.data == [] + assert result.pagination.total == 0 + assert result.pagination.total_pages == 0 def test_page_beyond_total(self): - items = list(range(25)) - params = PaginationParams(page=10, page_size=20) - result = paginate(items, params) - assert result.data == [] - assert result.pagination.total == 25 - assert result.pagination.total_pages == 2 - - def test_page_size_larger_than_total(self): + """页码超出总数""" items = list(range(5)) - params = PaginationParams(page=1, page_size=20) + params = PaginationParams(page=10, page_size=10) + result = paginate(items, params) - assert result.data == items + + assert result.data == [] + assert result.pagination.total == 5 assert result.pagination.total_pages == 1 - assert result.pagination.has_next is False + + def test_custom_page_size(self): + """自定义每页数量""" + items = list(range(100)) + params = PaginationParams(page=1, page_size=50) + + result = paginate(items, params) + + assert len(result.data) == 50 + assert result.pagination.total_pages == 2 def test_single_item(self): - items = [42] - params = PaginationParams(page=1, page_size=20) + """单条数据""" + items = ["only_one"] + params = PaginationParams(page=1, page_size=10) + result = paginate(items, params) - assert result.data == [42] + + assert result.data == ["only_one"] assert result.pagination.total == 1 + assert result.pagination.total_pages == 1 + + def test_generic_type_preserved(self): + """泛型类型数据正确""" + items = [{"id": 1, "name": "a"}, {"id": 2, "name": "b"}] + params = PaginationParams(page=1, page_size=10) - def test_page_size_1(self): - items = list(range(5)) - params = PaginationParams(page=3, page_size=1) result = paginate(items, params) - assert result.data == [2] - assert result.pagination.total_pages == 5 - def test_exact_page_size(self): - items = list(range(40)) - params = PaginationParams(page=2, page_size=20) - result = paginate(items, params) - assert result.data == list(range(20, 40)) - assert result.pagination.total_pages == 2 - assert result.pagination.has_next is False - - def test_string_items(self): - items = ["a", "b", "c", "d", "e"] - params = PaginationParams(page=2, page_size=2) - result = paginate(items, params) - assert result.data == ["c", "d"] - assert result.pagination.total == 5 - - def test_does_not_mutate_original_list(self): - items = list(range(10)) - original = items.copy() - params = PaginationParams(page=1, page_size=3) - paginate(items, params) - assert items == original + assert len(result.data) == 2 + assert result.data[0]["id"] == 1 diff --git a/tests/unit/test_password_hasher.py b/tests/unit/test_password_hasher.py index 7fe9c6dff..bc3a4787e 100755 --- a/tests/unit/test_password_hasher.py +++ b/tests/unit/test_password_hasher.py @@ -1,269 +1,250 @@ -""" -密码哈希工具测试 -""" +"""密码哈希与验证器单元测试.""" + +from __future__ import annotations import pytest -from packages.application.auth.password_hasher import PasswordHasher, PasswordValidator +from packages.application.auth.password_hasher import ( + PasswordHasher, + PasswordValidator, + password_hasher, + password_validator, +) class TestPasswordHasher: - """密码哈希测试""" + """PasswordHasher 测试""" - @pytest.fixture - def hasher(self): - """创建密码哈希器""" - return PasswordHasher(rounds=4) # 测试用低 cost,加快速度 - - def test_hash_password(self, hasher): - """测试密码哈希""" - password = "MySecurePassword123" - hashed = hasher.hash_password(password) + def test_hash_password_returns_string(self): + """哈希密码返回非空字符串""" + hasher = PasswordHasher(rounds=4) + hashed = hasher.hash_password("TestPass1!") assert isinstance(hashed, str) assert len(hashed) > 0 - assert hashed != password # 哈希后不等于原文 - assert hashed.startswith("$2b$") # bcrypt 格式 + assert hashed.startswith("$2") # bcrypt hash 格式 - def test_hash_same_password_different_result(self, hasher): - """测试相同密码每次哈希结果不同(因为 salt 不同)""" - password = "MySecurePassword123" - hash1 = hasher.hash_password(password) - hash2 = hasher.hash_password(password) + def test_hash_password_different_salts(self): + """相同密码每次哈希结果不同(加盐)""" + hasher = PasswordHasher(rounds=4) - assert hash1 != hash2 # salt 不同,哈希不同 + h1 = hasher.hash_password("SamePass1!") + h2 = hasher.hash_password("SamePass1!") - def test_verify_correct_password(self, hasher): - """测试验证正确的密码""" - password = "MySecurePassword123" - hashed = hasher.hash_password(password) + assert h1 != h2 - assert hasher.verify_password(password, hashed) is True + def test_verify_correct_password(self): + """正确密码验证通过""" + hasher = PasswordHasher(rounds=4) + hashed = hasher.hash_password("Correct1!") - def test_verify_incorrect_password(self, hasher): - """测试验证错误的密码""" - password = "MySecurePassword123" - hashed = hasher.hash_password(password) + assert hasher.verify_password("Correct1!", hashed) is True - assert hasher.verify_password("WrongPassword", hashed) is False + def test_verify_wrong_password(self): + """错误密码验证失败""" + hasher = PasswordHasher(rounds=4) + hashed = hasher.hash_password("Right123!") - def test_verify_empty_password(self, hasher): - """测试空密码验证""" - hashed = hasher.hash_password("test") + assert hasher.verify_password("Wrong123!", hashed) is False - assert hasher.verify_password("", hashed) is False + def test_hash_empty_password_raises(self): + """空密码哈希抛出 ValueError""" + hasher = PasswordHasher(rounds=4) - def test_verify_empty_hash(self, hasher): - """测试空哈希验证""" - assert hasher.verify_password("test", "") is False - - def test_verify_invalid_hash(self, hasher): - """测试无效的哈希""" - assert hasher.verify_password("test", "invalid-hash") is False - - def test_hash_empty_password(self, hasher): - """测试哈希空密码应该失败""" with pytest.raises(ValueError, match="Password cannot be empty"): hasher.hash_password("") - def test_invalid_rounds(self): - """测试无效的 rounds 参数""" + def test_verify_empty_password_returns_false(self): + """空密码验证返回 False""" + hasher = PasswordHasher(rounds=4) + hashed = hasher.hash_password("TestPass1!") + + assert hasher.verify_password("", hashed) is False + + def test_verify_empty_hash_returns_false(self): + """空哈希验证返回 False""" + hasher = PasswordHasher(rounds=4) + + assert hasher.verify_password("TestPass1!", "") is False + + def test_verify_invalid_hash_format(self): + """无效格式的哈希验证返回 False(不抛异常)""" + hasher = PasswordHasher(rounds=4) + + assert hasher.verify_password("TestPass1!", "not_a_valid_hash") is False + + def test_needs_rehash_same_rounds(self): + """相同 rounds 不需要重新哈希""" + hasher = PasswordHasher(rounds=4) + hashed = hasher.hash_password("TestPass1!") + + assert hasher.needs_rehash(hashed) is False + + def test_needs_rehash_different_rounds(self): + """不同 rounds 需要重新哈希""" + hasher_low = PasswordHasher(rounds=4) + hasher_high = PasswordHasher(rounds=5) + + hashed = hasher_low.hash_password("TestPass1!") + assert hasher_high.needs_rehash(hashed) is True + + def test_needs_rehash_invalid_hash(self): + """无效哈希格式返回 False(不抛异常)""" + hasher = PasswordHasher(rounds=4) + + assert hasher.needs_rehash("invalid_hash") is False + + def test_rounds_too_low_raises(self): + """rounds 小于 4 抛出 ValueError""" with pytest.raises(ValueError, match="rounds must be between 4 and 31"): - PasswordHasher(rounds=2) + PasswordHasher(rounds=3) + def test_rounds_too_high_raises(self): + """rounds 大于 31 抛出 ValueError""" with pytest.raises(ValueError, match="rounds must be between 4 and 31"): - PasswordHasher(rounds=50) + PasswordHasher(rounds=32) - def test_unicode_password(self, hasher): - """测试 Unicode 密码""" - password = "密码123!@#" - hashed = hasher.hash_password(password) + def test_rounds_boundary_values(self): + """rounds 边界值 4 和 31 是合法的""" + hasher_low = PasswordHasher(rounds=4) + hasher_high = PasswordHasher(rounds=31) - assert hasher.verify_password(password, hashed) is True - assert hasher.verify_password("错误密码", hashed) is False + assert hasher_low.rounds == 4 + assert hasher_high.rounds == 31 + + def test_hash_and_verify_various_passwords(self): + """多种密码的哈希-验证往返""" + hasher = PasswordHasher(rounds=4) + passwords = [ + "Simple12", + "C0mpl3x!@#", + " spaces ", + "中文密码123", + "a" * 50, # 50字节,在72字节限制内 + "12345678", + ] + for pwd in passwords: + hashed = hasher.hash_password(pwd) + assert hasher.verify_password(pwd, hashed) + assert not hasher.verify_password(pwd + "x", hashed) class TestPasswordValidator: - """密码验证器测试""" + """PasswordValidator 测试""" - @pytest.fixture - def validator(self): - """创建密码验证器""" - return PasswordValidator( - min_length=8, - require_uppercase=True, - require_lowercase=True, - require_digit=True, - require_special=False, - ) + def test_strong_password_passes(self): + """强密码通过验证""" + validator = PasswordValidator() + valid, error = validator.validate("Str0ngP@ss") - def test_valid_password(self, validator): - """测试有效密码""" - valid, error = validator.validate("MyPassword123") assert valid is True assert error is None - def test_password_too_short(self, validator): - """测试密码太短""" - valid, error = validator.validate("Pass1") - assert valid is False - assert "at least 8 characters" in error - - def test_password_no_uppercase(self, validator): - """测试没有大写字母""" - valid, error = validator.validate("mypassword123") - assert valid is False - assert "uppercase letter" in error - - def test_password_no_lowercase(self, validator): - """测试没有小写字母""" - valid, error = validator.validate("MYPASSWORD123") - assert valid is False - assert "lowercase letter" in error - - def test_password_no_digit(self, validator): - """测试没有数字""" - valid, error = validator.validate("MyPassword") - assert valid is False - assert "digit" in error - - def test_password_with_special_chars(self): - """测试要求特殊字符""" - validator = PasswordValidator( - min_length=8, - require_uppercase=True, - require_lowercase=True, - require_digit=True, - require_special=True, - ) - - # 没有特殊字符 - valid, error = validator.validate("MyPassword123") - assert valid is False - assert "special character" in error - - # 有特殊字符 - valid, error = validator.validate("MyPassword123!") - assert valid is True - assert error is None - - def test_empty_password(self, validator): - """测试空密码""" + def test_empty_password_fails(self): + """空密码验证失败""" + validator = PasswordValidator() valid, error = validator.validate("") + assert valid is False - assert "cannot be empty" in error + assert "empty" in error.lower() + + def test_too_short_fails(self): + """密码太短失败""" + validator = PasswordValidator(min_length=8) + valid, error = validator.validate("Sh0rt!") + + assert valid is False + assert "at least 8" in error + + def test_no_uppercase_fails(self): + """没有大写字母失败""" + validator = PasswordValidator(require_uppercase=True) + valid, error = validator.validate("lowercase1!") + + assert valid is False + assert "uppercase" in error.lower() + + def test_no_lowercase_fails(self): + """没有小写字母失败""" + validator = PasswordValidator(require_lowercase=True) + valid, error = validator.validate("UPPERCASE1!") + + assert valid is False + assert "lowercase" in error.lower() + + def test_no_digit_fails(self): + """没有数字失败""" + validator = PasswordValidator(require_digit=True) + valid, error = validator.validate("NoDigitsHere!") + + assert valid is False + assert "digit" in error.lower() + + def test_no_special_not_required_passes(self): + """不要求特殊字符时,不含特殊字符也通过""" + validator = PasswordValidator(require_special=False) + valid, error = validator.validate("NoSpecial1") + + assert valid is True + + def test_no_special_required_fails(self): + """要求特殊字符时,不含特殊字符失败""" + validator = PasswordValidator(require_special=True) + valid, error = validator.validate("NoSpecial1") + + assert valid is False + assert "special" in error.lower() def test_custom_min_length(self): - """测试自定义最小长度""" + """自定义最小长度""" validator = PasswordValidator( min_length=12, require_uppercase=False, require_lowercase=False, require_digit=False, + ) + valid, _ = validator.validate("123456789012") # 12字符 + assert valid is True + + valid, _ = validator.validate("12345678901") # 11字符 + assert valid is False + + def test_all_requirements_disabled(self): + """所有要求都禁用时,任意非空密码都通过""" + validator = PasswordValidator( + min_length=1, + require_uppercase=False, + require_lowercase=False, + require_digit=False, require_special=False, ) + valid, error = validator.validate("x") - valid, error = validator.validate("short") - assert valid is False - assert "at least 12 characters" in error - - valid, error = validator.validate("longenoughpassword") assert valid is True assert error is None + def test_special_characters_recognized(self): + """各种特殊字符都被识别""" + validator = PasswordValidator(require_special=True, require_uppercase=False, require_lowercase=False) + specials = ["!", "@", "#", "$", "%", "^", "&", "*", "(", ")", "-", "_", "=", "+"] + for ch in specials: + valid, _ = validator.validate(f"abcd1234{ch}") + assert valid is True, f"Special char '{ch}' not recognized" -class TestPasswordHandler: - """Password Handler 委托层测试""" - def test_hash_and_verify_password(self): - """测试哈希和验证密码""" - from packages.application.auth.password_handler import PasswordHandler +class TestGlobalInstances: + """全局实例测试""" - handler = PasswordHandler(rounds=4) - hashed = handler.hash_password("MySecurePass123") + def test_global_password_hasher_exists(self): + """全局 password_hasher 实例存在""" + assert password_hasher is not None + assert isinstance(password_hasher, PasswordHasher) + assert password_hasher.rounds == 12 - assert hashed != "MySecurePass123" - assert len(hashed) > 20 - assert handler.verify_password("MySecurePass123", hashed) is True - assert handler.verify_password("WrongPassword", hashed) is False - - def test_hash_empty_password_raises(self): - """测试空密码抛出异常""" - from packages.application.auth.password_handler import PasswordHandler - - handler = PasswordHandler(rounds=4) - with pytest.raises(ValueError): - handler.hash_password("") - - def test_needs_rehash(self): - """测试检测需要重新哈希""" - from packages.application.auth.password_handler import PasswordHandler - - handler = PasswordHandler(rounds=4) - hashed = handler.hash_password("TestPass123") - - # 相同 rounds 不需要重新哈希 - assert handler.needs_rehash(hashed) is False - - # 用更高 rounds 的 handler 检查,应该需要重新哈希 - # 注意:bcrypt 的 rounds 体现在 hash 中,这里用不同 rounds 测试 - high_rounds_handler = PasswordHandler(rounds=5) - # 低 rounds 的 hash 在高 rounds 配置下应该需要 rehash - assert high_rounds_handler.needs_rehash(hashed) is True - - def test_validate_strength(self): - """测试密码强度验证""" - from packages.application.auth.password_handler import PasswordHandler - - handler = PasswordHandler(rounds=4) - - # 弱密码 - valid, error = handler.validate_strength("weak") - assert valid is False - assert error is not None - - # 强密码 - valid, error = handler.validate_strength("StrongPass123") - assert valid is True - assert error is None - - def test_configure_and_get_default_handler(self): - """测试配置和获取全局默认 handler""" - from packages.application.auth import password_handler as handler_module - from packages.application.auth.password_handler import ( - configure_password_handler, - get_password_handler, - ) - - # 重置全局状态 - handler_module._default_handler = None - - # 配置 - handler = configure_password_handler(rounds=4) - assert handler is not None - - # 获取 - same_handler = get_password_handler() - assert same_handler is handler - - # 验证能正常工作 - hashed = same_handler.hash_password("TestPass123") - assert same_handler.verify_password("TestPass123", hashed) is True - - # 重置全局状态,避免影响其他测试 - handler_module._default_handler = None - - def test_get_password_handler_auto_creates_default(self): - """测试未配置时获取 handler 会自动创建默认实例""" - from packages.application.auth import password_handler as handler_module - from packages.application.auth.password_handler import get_password_handler - - # 重置全局状态 - handler_module._default_handler = None - - # 自动创建默认实例 - handler = get_password_handler() - assert handler is not None - - # 重置 - handler_module._default_handler = None + def test_global_password_validator_exists(self): + """全局 password_validator 实例存在""" + assert password_validator is not None + assert isinstance(password_validator, PasswordValidator) + assert password_validator.min_length == 8 + assert password_validator.require_uppercase is True + assert password_validator.require_special is False diff --git a/tests/unit/test_text_splitter.py b/tests/unit/test_text_splitter.py index 81fc33a06..d52bec9a0 100755 --- a/tests/unit/test_text_splitter.py +++ b/tests/unit/test_text_splitter.py @@ -1,315 +1,147 @@ -""" -text_splitter 长文本分段工具单元测试 +"""文本分段工具单元测试.""" -覆盖: -- 空文本 / 短文本 -- 句子边界分段(。!?;\n . ! ? ;) -- 超长句子硬切 -- 过短段落合并 -- max_chars 参数 -- 中英文混合 -""" +from __future__ import annotations import pytest from packages.application.tts_job.text_splitter import split_text -# ============================================================ -# 基础场景 -# ============================================================ +class TestSplitText: + """split_text 函数测试""" -class TestBasicCases: - """基础场景""" - - def test_empty_text_returns_empty_list(self): + def test_empty_string_returns_empty_list(self): + """空字符串返回空列表""" assert split_text("") == [] - def test_whitespace_only_returns_empty(self): - assert split_text(" \n\n ") == [] + def test_whitespace_only_returns_empty_list(self): + """纯空白字符返回空列表""" + assert split_text(" \n \t ") == [] - def test_short_text_single_segment(self): + def test_short_text_returns_single_segment(self): + """短文本直接返回单段""" text = "这是一段短文本。" result = split_text(text, max_chars=500) assert result == [text] - def test_exactly_max_chars_single_segment(self): - text = "a" * 500 - result = split_text(text, max_chars=500) + def test_text_length_equals_max_chars(self): + """文本长度恰好等于 max_chars 时返回单段""" + text = "a" * 100 + result = split_text(text, max_chars=100) assert len(result) == 1 - assert len(result[0]) == 500 + assert len(result[0]) == 100 - def test_text_stripped(self): - text = " 你好世界。 " - result = split_text(text, max_chars=500) - assert result == ["你好世界。"] + def test_splits_on_sentence_boundary(self): + """在句子边界处分段""" + # 构造长文本,确保超过 max_chars + sentences = ["今天天气真好。我们一起去公园散步吧。", "公园里有很多花。还有很多小朋友在玩耍。"] * 10 + text = "".join(sentences) + result = split_text(text, max_chars=200) -# ============================================================ -# 句子边界分段 -# ============================================================ - - -class TestSentenceBoundarySplitting: - """句子边界分段""" - - def test_split_by_chinese_period(self): - text = "第一句。第二句。第三句。" - # 三句都很短,应该合并成一段 - result = split_text(text, max_chars=500) - assert len(result) == 1 - - def test_split_by_chinese_period_long_text(self): - """多段长句子,按句号分段""" - sentence1 = "我是第一句" + "啊" * 100 + "。" - sentence2 = "我是第二句" + "哦" * 100 + "。" - sentence3 = "我是第三句" + "嗯" * 100 + "。" - text = sentence1 + sentence2 + sentence3 - - result = split_text(text, max_chars=150) - # 每句106字符,超过150的阈值?不,106<150 - # 但累计到一定程度会切 assert len(result) >= 2 # 每段都不超过 max_chars - for seg in result: - assert len(seg) <= 150 - - def test_split_by_question_mark(self): - text = "你是谁?你从哪里来?你要到哪里去?" - result = split_text(text, max_chars=500) - # 三句都很短,合并成一段 - assert len(result) == 1 - - def test_split_by_exclamation_mark(self): - text = "太棒了!太厉害了!太牛了!" - result = split_text(text, max_chars=500) - assert len(result) == 1 - - def test_split_by_newline(self): - text = "第一段\n第二段\n第三段" - result = split_text(text, max_chars=500) - assert len(result) == 1 - - def test_split_by_semicolon(self): - text = "第一部分;第二部分;第三部分。" - result = split_text(text, max_chars=500) - assert len(result) == 1 - - def test_mixed_punctuation(self): - """混合标点符号的句子边界""" - parts = [] - for i in range(20): - parts.append(f"第{i}句的内容" + "字" * 30 + "。") - text = "".join(parts) - - result = split_text(text, max_chars=200) - # 每句约35字符,200字符大约能放5-6句 - assert len(result) >= 2 for seg in result: assert len(seg) <= 200 - def test_english_period_splitting(self): - text = "Hello. How are you. I am fine." - result = split_text(text, max_chars=500) - assert len(result) == 1 + def test_all_segments_within_max_chars(self): + """所有分段都不超过 max_chars""" + text = "这是第一句话。这是第二句话。这是第三句话。这是第四句话。这是第五句话。" * 10 - def test_english_question(self): - text = "What? Why? How?" - result = split_text(text, max_chars=500) - assert len(result) == 1 - - -# ============================================================ -# 超长硬切 -# ============================================================ - - -class TestLongSentenceHardCut: - """超长句子硬切""" - - def test_single_very_long_sentence_hard_cut(self): - """单个超长句子,没有标点,硬切""" - text = "字" * 1000 - result = split_text(text, max_chars=500) - assert len(result) == 2 - assert len(result[0]) == 500 - assert len(result[1]) == 500 - - def test_three_times_max_chars(self): - text = "字" * 1500 - result = split_text(text, max_chars=500) - assert len(result) == 3 - for seg in result: - assert len(seg) == 500 - - def test_not_exact_multiple(self): - text = "字" * 1250 - result = split_text(text, max_chars=500) - assert len(result) == 3 - assert len(result[0]) == 500 - assert len(result[1]) == 500 - assert len(result[2]) == 250 - - def test_all_segments_within_limit(self): - """所有段都不超过 max_chars""" - import random - - random.seed(42) - # 生成随机长度的文本 - text = "".join(random.choices("字字字字。!?;\n", k=5000)) - for max_chars in [100, 200, 500]: - result = split_text(text, max_chars=max_chars) - for i, seg in enumerate(result): - assert len(seg) <= max_chars, f"Segment {i} length {len(seg)} > {max_chars}" - - -# ============================================================ -# 过短段落合并 -# ============================================================ - - -class TestShortSegmentMerging: - """过短段落合并""" - - def test_short_final_segment_merged(self): - """最后一段过短,应该合并到前一段""" - # 构造:前一段接近上限,后一段很短 - long_part = "字" * 480 + "。" - short_part = "好的。" - text = long_part + short_part - - result = split_text(text, max_chars=500) - # 两段加起来 481+3=484 < 500,可能合并 - # 但要看具体实现... - # 至少验证所有段不超长 - for seg in result: - assert len(seg) <= 500 - - def test_multiple_short_segments(self): - """多个短段落应该合并""" - sentences = ["你好。", "我好。", "大家好。", "今天天气不错。", "适合出去玩。"] - text = "".join(sentences) - result = split_text(text, max_chars=500) - # 5个短句子,应该合并成一段 - assert len(result) == 1 - - -# ============================================================ -# max_chars 参数 -# ============================================================ - - -class TestMaxCharsParameter: - """max_chars 参数""" - - def test_small_max_chars(self): - text = "一二三四五六七八九十一二三四五六七八九十。" - result = split_text(text, max_chars=10) - # 应该被切成多段 - assert len(result) >= 2 - for seg in result: - assert len(seg) <= 10 - - def test_custom_max_chars_200(self): - text = "测试文本" * 100 # 400字符 - result = split_text(text, max_chars=200) - assert len(result) == 2 - assert len(result[0]) == 200 - assert len(result[1]) == 200 - - def test_very_small_max_chars(self): - text = "abcdefghij" - result = split_text(text, max_chars=3) - assert len(result) >= 3 - for seg in result: - assert len(seg) <= 3 - - -# ============================================================ -# 中英文混合 -# ============================================================ - - -class TestMixedContent: - """中英文混合内容""" - - def test_chinese_english_mixed(self): - text = "今天天气很好,Today is sunny. 我们去公园玩吧!Let's go to the park." - result = split_text(text, max_chars=500) - assert len(result) == 1 - assert result[0] == text.strip() - - def test_mixed_long_text(self): - parts = [] - for i in range(50): - parts.append(f"第{i}段中文内容" + "字" * 20 + ". English part " + "word " * 10 + "。") - text = "".join(parts) - - result = split_text(text, max_chars=300) - assert len(result) >= 2 - for seg in result: - assert len(seg) <= 300 - - -# ============================================================ -# 输出完整性 -# ============================================================ - - -class TestOutputIntegrity: - """输出完整性验证""" - - def test_combined_length_equals_original(self): - """所有段拼接起来(去掉空段)应该等于原文长度""" - text = "这是第一段。这是第二段。这是第三段。这是第四段。这是第五段。" * 20 result = split_text(text, max_chars=100) - combined = "".join(result) - # 由于 strip 可能去掉一些空格,原文也 strip 比较 - assert len(combined) == len(text.strip()) - def test_order_preserved(self): - """分段后再拼接,文本顺序不变""" - text = "第一。第二。第三。第四。第五。" * 10 - result = split_text(text, max_chars=50) - combined = "".join(result) - assert combined == text.strip() - - def test_no_empty_strings_in_result(self): - """结果中没有空字符串""" - text = "句子一。句子二。句子三。" - result = split_text(text, max_chars=10) for seg in result: - assert seg != "" - assert len(seg) > 0 + assert len(seg) <= 100 + def test_long_single_sentence_hard_cut(self): + """超长单句会被硬切""" + text = "a" * 1000 # 没有标点 -# ============================================================ -# 边界情况 -# ============================================================ - - -class TestEdgeCases: - """边界情况""" - - def test_single_character(self): - assert split_text("一", max_chars=500) == ["一"] - - def test_only_punctuation(self): - text = "。。。。。" - result = split_text(text, max_chars=500) - # 都是标点,也算文本 - assert len(result) == 1 - - def test_only_newlines(self): - text = "\n\n\n" - result = split_text(text, max_chars=500) - assert result == [] - - def test_long_text_many_sentences(self): - """大量句子的长文本""" - sentences = [f"第{i}句的完整内容。" for i in range(100)] - text = "".join(sentences) result = split_text(text, max_chars=200) - assert len(result) >= 5 + + assert len(result) > 1 for seg in result: assert len(seg) <= 200 + + def test_newline_is_sentence_end(self): + """换行符作为句子结束符""" + text = "第一行内容\n第二行内容\n第三行内容" * 10 + + result = split_text(text, max_chars=50) + + assert len(result) > 1 + for seg in result: + assert len(seg) <= 50 + + def test_chinese_punctuation(self): + """中文标点(。!?;)作为句子结束符""" + text = "你好!今天吃什么?我吃米饭;你呢?我也吃米饭。" * 10 + + result = split_text(text, max_chars=80) + + for seg in result: + assert len(seg) <= 80 + + def test_english_punctuation(self): + """英文标点(.!?;)作为句子结束符""" + text = "Hello! How are you? I'm fine; thank you. Good bye." * 10 + + result = split_text(text, max_chars=80) + + for seg in result: + assert len(seg) <= 80 + + def test_merged_short_segments(self): + """过短的段落会被合并""" + # 构造很多短句 + text = "你好。再见。谢谢。抱歉。好的。不行。可以。去吧。" * 5 # 每句3-4字 + + result = split_text(text, max_chars=100) + + # 合并后段数应该比单纯按句切的少 + assert len(result) < len(text) // 3 # 粗略估计 + for seg in result: + assert len(seg) <= 100 + + def test_preserves_content(self): + """分段后内容总和与原文基本一致(忽略strip的空白)""" + text = "这是测试文本。包含多个句子。用来验证分段正确性。" * 5 + + result = split_text(text, max_chars=50) + + # 合并所有分段,去掉空白后应该与原文去掉空白后基本一致 + combined = "".join(result).replace(" ", "") + original = text.strip().replace(" ", "") + assert combined == original + + def test_custom_max_chars(self): + """支持自定义 max_chars""" + text = "测试" * 100 # 200字 + + result_50 = split_text(text, max_chars=50) + result_100 = split_text(text, max_chars=100) + + # max_chars 越小,段数应该越多 + assert len(result_50) >= len(result_100) + + def test_single_char_text(self): + """单字符文本""" + assert split_text("好", max_chars=10) == ["好"] + + def test_text_with_only_punctuation(self): + """纯标点文本""" + text = "。。。。。。。。。。" # 10个句号 + result = split_text(text, max_chars=5) + + assert len(result) >= 1 + for seg in result: + assert len(seg) <= 5 + + def test_mixed_content(self): + """中英文混合内容""" + text = "今天的天气是 sunny and warm。我们去了 park 玩。真的很开心!" * 5 + + result = split_text(text, max_chars=80) + + for seg in result: + assert len(seg) <= 80