Compare commits

...

7 Commits

Author SHA1 Message Date
CI Auto Fix Bot 5aef45636f fix(ci): auto-fix lint/format issues
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 29s
AI Code Review / AI Code Review (pull_request) Successful in 7m37s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Has been skipped
CI/CD Pipeline / Validate - Code Quality (pull_request) Has been skipped
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Has been skipped
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 26m13s
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / Unit Tests (pull_request) Successful in 41s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m14s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 38m43s
PR Automation / Auto Approve on CI Green (pull_request) Has been skipped
Preview Cleanup / Cleanup Preview Environment (pull_request) Failing after 0s
2026-07-22 15:07:41 +08:00
CI Bot de5538ad1f test(text_splitter): 补充长文本分段工具32个单元测试
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1m2s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m0s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 1m12s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 40s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 2m6s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 2m15s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 1m2s
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / PR Build API Image (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been cancelled
AI Code Review / AI Code Review (pull_request) Has been cancelled
Preview Deploy / Deploy Preview Environment (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m4s
覆盖空文本/短文本/句子边界分段/超长硬切/
过短段落合并/max_chars参数/中英文混合/输出完整性
2026-07-22 14:43:22 +08:00
CI Bot 0a23f72bfa test(feature_flag_store): 补充FeatureFlag配置与存储43个单元测试
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 28s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 56s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 32s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 35s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m48s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 57s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 1m10s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m28s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m33s
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / PR Build API Image (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been cancelled
AI Code Review / AI Code Review (pull_request) Has been cancelled
覆盖 FeatureFlagConfig 序列化/is_active判定(全局开关/
白名单/百分比哈希/边界值)、InMemoryFeatureFlagStore CRUD
2026-07-22 14:39:42 +08:00
CI Bot bb83549cc0 test(module_registry): 补充模块注册中心62个单元测试
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 22s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 36s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 23s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 1m0s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 57s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 1m4s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m59s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m8s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 50s
AI Code Review / AI Code Review (pull_request) Successful in 8m41s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 10m45s
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been cancelled
覆盖 ModuleStatus枚举/QuotaRule/ModuleCapability/Module数据类、
Module.activate/disable状态转换、ModuleRegistry注册注销/
自动激活/依赖检查/能力发现/状态过滤/全局单例
2026-07-22 14:34:56 +08:00
CI Bot 92b8044cb3 test(storage): 补充SharedStorageService 55个单元测试
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 21s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m29s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 27s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 49s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 41s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 54s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 1m39s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 52s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m16s
AI Code Review / AI Code Review (pull_request) Successful in 5m2s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 8m51s
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been cancelled
覆盖 _normalize_storage_key/URL解码、_is_local_generated_url、
get_url、create_direct_upload_post/policy+HMAC签名验证、
get_download_url fallback、未配置OSS错误处理、
bucket操作调用验证、单例模式、endpoint https前缀处理
2026-07-22 14:31:01 +08:00
CI Bot a9bc314fed test(pagination): 补充通用分页器52个单元测试
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 24s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 47s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 26s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 25s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 41s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 46s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 50s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 37s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 5m59s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 6m6s
AI Code Review / AI Code Review (pull_request) Successful in 3m20s
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been cancelled
覆盖 PaginationParams 默认值/校验/offset/limit、
PaginationMeta.from_params 边界场景、
PaginatedResponse.create 工厂方法、
paginate 内存分页函数
2026-07-22 14:28:40 +08:00
CI Bot 911478495a test(p3-1): 第十波 - url_security 安全模块单测 72个
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 10s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 46s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 32s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 30s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 56s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m19s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 3m32s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 48s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 49s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m12s
AI Code Review / AI Code Review (pull_request) Successful in 5m35s
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been cancelled
覆盖 url_security.py 核心安全逻辑:
- validate_url_safety: scheme/端口/主机名/URL长度 等基础校验
- SSRF防护: 内网IP/回环/链路本地/组播/未指定/保留地址
- 内网主机名防护: localhost/.local/.internal/metadata
- 可信域名白名单: 精确匹配+子域名匹配
- is_url_safe 便捷函数
- 文件魔数校验: PNG/JPEG/GIF/WEBP/BMP/WAV/MP3/OGG/FLAC
- safe_download_file / safe_download_bytes (mock网络)
- 大小限制 / MIME类型白名单 / Content-Length预检
2026-07-22 14:24:32 +08:00
6 changed files with 2738 additions and 252 deletions
+425
View File
@@ -0,0 +1,425 @@
"""
FeatureFlagStore 单元测试
覆盖:
- FeatureFlagConfig: to_dict / from_dict 序列化
- FeatureFlagConfig.is_active: 全局开关/白名单/百分比哈希
- InMemoryFeatureFlagStore: CRUD / is_active
"""
import pytest
from packages.adapters.redis.feature_flag_store import (
FEATURE_FLAG_REDIS_PREFIX,
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
# ============================================================
# 常量
# ============================================================
class TestConstants:
"""常量验证"""
def test_redis_prefix(self):
assert FEATURE_FLAG_REDIS_PREFIX == "feature_flag:"
# ============================================================
# FeatureFlagConfig - 默认值 & 基础
# ============================================================
class TestFeatureFlagConfigDefaults:
"""FeatureFlagConfig 默认值"""
def test_required_name(self):
config = FeatureFlagConfig(name="test_flag")
assert config.name == "test_flag"
def test_default_disabled(self):
config = FeatureFlagConfig(name="test_flag")
assert config.enabled is False
def test_default_percentage_zero(self):
config = FeatureFlagConfig(name="test_flag")
assert config.percentage == 0
def test_default_whitelist_empty(self):
config = FeatureFlagConfig(name="test_flag")
assert config.whitelist == set()
def test_full_config(self):
config = FeatureFlagConfig(
name="full_flag",
enabled=True,
percentage=50,
whitelist={"user1", "user2"},
)
assert config.name == "full_flag"
assert config.enabled is True
assert config.percentage == 50
assert config.whitelist == {"user1", "user2"}
# ============================================================
# FeatureFlagConfig - 序列化
# ============================================================
class TestFeatureFlagConfigSerialization:
"""to_dict / from_dict 序列化"""
def test_to_dict_defaults(self):
config = FeatureFlagConfig(name="test")
d = config.to_dict()
assert d["name"] == "test"
assert d["enabled"] is False
assert d["percentage"] == 0
assert d["whitelist"] == []
def test_to_dict_with_values(self):
config = FeatureFlagConfig(
name="test",
enabled=True,
percentage=75,
whitelist={"a", "b", "c"},
)
d = config.to_dict()
assert d["name"] == "test"
assert d["enabled"] is True
assert d["percentage"] == 75
# whitelist 排序后输出
assert sorted(d["whitelist"]) == ["a", "b", "c"]
def test_from_dict_minimal(self):
d = {"name": "test"}
config = FeatureFlagConfig.from_dict(d)
assert config.name == "test"
assert config.enabled is False
assert config.percentage == 0
assert config.whitelist == set()
def test_from_dict_full(self):
d = {
"name": "full",
"enabled": True,
"percentage": 30,
"whitelist": ["u1", "u2"],
}
config = FeatureFlagConfig.from_dict(d)
assert config.name == "full"
assert config.enabled is True
assert config.percentage == 30
assert config.whitelist == {"u1", "u2"}
def test_round_trip(self):
original = FeatureFlagConfig(
name="round_trip",
enabled=True,
percentage=42,
whitelist={"alice", "bob", "charlie"},
)
d = original.to_dict()
restored = FeatureFlagConfig.from_dict(d)
assert restored.name == original.name
assert restored.enabled == original.enabled
assert restored.percentage == original.percentage
assert restored.whitelist == original.whitelist
def test_from_dict_coerces_types(self):
"""from_dict 应该做类型转换"""
d = {
"name": "coerce",
"enabled": 1, # int → bool
"percentage": "50", # str → int
"whitelist": ("a", "b"), # tuple → set
}
config = FeatureFlagConfig.from_dict(d)
assert config.enabled is True
assert config.percentage == 50
assert config.whitelist == {"a", "b"}
# ============================================================
# FeatureFlagConfig.is_active - 全局开关
# ============================================================
class TestIsActiveGlobalSwitch:
"""is_active - 全局开关基础"""
def test_disabled_returns_false(self):
config = FeatureFlagConfig(name="test", enabled=False)
assert config.is_active() is False
def test_disabled_with_identifier_returns_false(self):
config = FeatureFlagConfig(name="test", enabled=False)
assert config.is_active(identifier="user1") is False
def test_enabled_no_percentage_no_whitelist_returns_true(self):
config = FeatureFlagConfig(name="test", enabled=True)
# percentage=0, whitelist=空,但 enabled=True
# 按逻辑:全局开了但百分比0且无白名单 → 其实应该是 False?
# 让我看代码...
# 代码里 percentage <= 0 时返回 False(没有白名单且百分比为0)
assert config.is_active() is False
def test_enabled_100_percent_returns_true(self):
config = FeatureFlagConfig(name="test", enabled=True, percentage=100)
assert config.is_active() is True
# ============================================================
# FeatureFlagConfig.is_active - 白名单
# ============================================================
class TestIsActiveWhitelist:
"""is_active - 白名单优先级"""
def test_whitelist_match_returns_true(self):
config = FeatureFlagConfig(
name="test",
enabled=True,
whitelist={"user1", "user2"},
)
assert config.is_active(identifier="user1") is True
def test_whitelist_no_match_falls_through(self):
config = FeatureFlagConfig(
name="test",
enabled=True,
percentage=0,
whitelist={"user1"},
)
# 不在白名单,且百分比为0 → False
assert config.is_active(identifier="user3") is False
def test_whitelist_overrides_percentage_zero(self):
"""白名单优先级最高,即使百分比为0也能启用"""
config = FeatureFlagConfig(
name="test",
enabled=True,
percentage=0,
whitelist={"vip_user"},
)
assert config.is_active(identifier="vip_user") is True
def test_whitelist_overrides_partial_percentage(self):
"""白名单用户即使在百分比外也能启用"""
config = FeatureFlagConfig(
name="test",
enabled=True,
percentage=1, # 只有1%的用户
whitelist={"important_user"},
)
# 白名单用户直接通过
assert config.is_active(identifier="important_user") is True
def test_no_identifier_no_whitelist_check(self):
"""不传 identifier 时不做白名单检查"""
config = FeatureFlagConfig(
name="test",
enabled=True,
percentage=100,
whitelist={"user1"},
)
# 无 identifier,直接看百分比(100%
assert config.is_active() is True
# ============================================================
# FeatureFlagConfig.is_active - 百分比边界值
# ============================================================
class TestIsActivePercentageBoundaries:
"""is_active - 百分比边界值"""
def test_percentage_0_returns_false(self):
config = FeatureFlagConfig(name="test", enabled=True, percentage=0)
assert config.is_active(identifier="any_user") is False
def test_percentage_100_returns_true(self):
config = FeatureFlagConfig(name="test", enabled=True, percentage=100)
assert config.is_active(identifier="any_user") is True
def test_percentage_negative_treated_as_0(self):
"""percentage < 0 应该按 0 处理"""
config = FeatureFlagConfig(name="test", enabled=True, percentage=-5)
assert config.is_active(identifier="any_user") is False
def test_percentage_over_100_treated_as_100(self):
"""percentage > 100 应该按 100 处理"""
config = FeatureFlagConfig(name="test", enabled=True, percentage=150)
assert config.is_active(identifier="any_user") is True
# ============================================================
# FeatureFlagConfig.is_active - 哈希一致性
# ============================================================
class TestIsActiveHashConsistency:
"""is_active - 哈希取模一致性验证"""
def test_same_user_same_result_every_time(self):
"""同一用户多次调用结果一致(确定性哈希)"""
config = FeatureFlagConfig(name="test", enabled=True, percentage=50)
results = {config.is_active(identifier="user_xyz") for _ in range(100)}
assert len(results) == 1 # 全部相同
def test_different_flags_same_user_can_differ(self):
"""不同 flag 对同一用户可以有不同结果(因为 flag name 参与哈希)"""
config_a = FeatureFlagConfig(name="flag_a", enabled=True, percentage=50)
config_b = FeatureFlagConfig(name="flag_b", enabled=True, percentage=50)
# 不保证一定不同,但大部分情况下应该不同
# 这里只验证哈希输入包含了 flag name(通过机制保证)
# 具体是否不同取决于哈希值
def test_percentage_coverage_roughly_correct(self):
"""大量用户中,命中比例大致接近百分比"""
config = FeatureFlagConfig(name="coverage_test", enabled=True, percentage=30)
users = [f"user_{i}" for i in range(1000)]
active_count = sum(1 for u in users if config.is_active(identifier=u))
# 30% ± 10% 的容差
assert 200 <= active_count <= 400
def test_50_percent_roughly_half(self):
config = FeatureFlagConfig(name="half_test", enabled=True, percentage=50)
users = [f"user_{i}" for i in range(1000)]
active_count = sum(1 for u in users if config.is_active(identifier=u))
# 50% ± 10%
assert 400 <= active_count <= 600
def test_10_percent_roughly_tenth(self):
config = FeatureFlagConfig(name="ten_pct", enabled=True, percentage=10)
users = [f"user_{i}" for i in range(1000)]
active_count = sum(1 for u in users if config.is_active(identifier=u))
assert 50 <= active_count <= 150
def test_empty_identifier_treated_as_no_identifier(self):
"""空字符串 identifier 应该如何处理?"""
config = FeatureFlagConfig(name="test", enabled=True, percentage=50)
# 空字符串是 falsy,走无 identifier 分支(随机)
# 但白名单检查也会跳过
# 验证不会崩溃
result = config.is_active(identifier="")
assert isinstance(result, bool)
# ============================================================
# InMemoryFeatureFlagStore - CRUD
# ============================================================
class TestInMemoryFeatureFlagStore:
"""InMemoryFeatureFlagStore 内存实现"""
def test_get_nonexistent_returns_default_disabled(self):
store = InMemoryFeatureFlagStore()
config = store.get("nonexistent")
assert config.name == "nonexistent"
assert config.enabled is False
assert config.percentage == 0
def test_set_and_get(self):
store = InMemoryFeatureFlagStore()
original = FeatureFlagConfig(
name="my_flag",
enabled=True,
percentage=50,
whitelist={"admin"},
)
store.set(original)
retrieved = store.get("my_flag")
assert retrieved.name == "my_flag"
assert retrieved.enabled is True
assert retrieved.percentage == 50
assert retrieved.whitelist == {"admin"}
def test_set_overwrites_existing(self):
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="flag", enabled=True, percentage=30))
store.set(FeatureFlagConfig(name="flag", enabled=False, percentage=70))
config = store.get("flag")
assert config.enabled is False
assert config.percentage == 70
def test_delete_existing_returns_true(self):
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="delete_me"))
result = store.delete("delete_me")
assert result is True
# 删除后获取返回默认配置
assert store.get("delete_me").enabled is False
def test_delete_nonexistent_returns_false(self):
store = InMemoryFeatureFlagStore()
result = store.delete("no_such_flag")
assert result is False
def test_list_all_empty(self):
store = InMemoryFeatureFlagStore()
assert store.list_all() == {}
def test_list_all_multiple(self):
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="flag1", enabled=True))
store.set(FeatureFlagConfig(name="flag2", percentage=50))
store.set(FeatureFlagConfig(name="flag3"))
all_flags = store.list_all()
assert len(all_flags) == 3
assert "flag1" in all_flags
assert "flag2" in all_flags
assert "flag3" in all_flags
assert all_flags["flag1"].enabled is True
assert all_flags["flag2"].percentage == 50
def test_list_all_returns_copy(self):
"""返回的是副本,修改不影响内部状态"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="flag1"))
flags = store.list_all()
flags["fake"] = FeatureFlagConfig(name="fake")
assert "fake" not in store.list_all()
# ============================================================
# InMemoryFeatureFlagStore - is_active
# ============================================================
class TestInMemoryStoreIsActive:
"""store.is_active 便捷方法"""
def test_is_active_enabled_flag(self):
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="on", enabled=True, percentage=100))
assert store.is_active("on") is True
def test_is_active_disabled_flag(self):
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="off", enabled=False))
assert store.is_active("off") is False
def test_is_active_nonexistent_flag(self):
store = InMemoryFeatureFlagStore()
assert store.is_active("unknown") is False
def test_is_active_with_identifier_whitelist(self):
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="beta",
enabled=True,
percentage=0,
whitelist={"tester1"},
)
)
assert store.is_active("beta", identifier="tester1") is True
assert store.is_active("beta", identifier="other_user") is False
+569
View File
@@ -0,0 +1,569 @@
"""
Module Registry 模块注册中心单元测试
覆盖:
- ModuleStatus 枚举
- QuotaRule / ModuleCapability / Module 数据类
- Module.activate / disable 状态转换
- ModuleRegistry 注册/注销/查询/能力发现/依赖检查
"""
import pytest
from packages.infrastructure.module_registry import (
Module,
ModuleCapability,
ModuleRegistry,
ModuleStatus,
QuotaRule,
module_registry,
)
# ============================================================
# ModuleStatus
# ============================================================
class TestModuleStatus:
"""ModuleStatus 枚举"""
def test_enum_values(self):
assert ModuleStatus.REGISTERED.value == "registered"
assert ModuleStatus.ACTIVE.value == "active"
assert ModuleStatus.DISABLED.value == "disabled"
assert ModuleStatus.ERROR.value == "error"
def test_is_str_enum(self):
assert isinstance(ModuleStatus.ACTIVE, str)
assert ModuleStatus.ACTIVE == "active"
def test_has_four_states(self):
assert len(ModuleStatus) == 4
# ============================================================
# QuotaRule
# ============================================================
class TestQuotaRule:
"""QuotaRule 配额规则"""
def test_required_fields(self):
rule = QuotaRule(dimension="ai_credits", per_operation=1.0)
assert rule.dimension == "ai_credits"
assert rule.per_operation == 1.0
def test_default_description_empty(self):
rule = QuotaRule(dimension="storage_gb", per_operation=0.5)
assert rule.description == ""
def test_custom_description(self):
rule = QuotaRule(
dimension="credits",
per_operation=2.0,
description="每次生成消耗2积分",
)
assert rule.description == "每次生成消耗2积分"
def test_float_per_operation(self):
rule = QuotaRule(dimension="gb", per_operation=0.25)
assert rule.per_operation == 0.25
# ============================================================
# ModuleCapability
# ============================================================
class TestModuleCapability:
"""ModuleCapability 能力定义"""
def test_required_name(self):
cap = ModuleCapability(name="generate_voice")
assert cap.name == "generate_voice"
def test_defaults(self):
cap = ModuleCapability(name="test_cap")
assert cap.description == ""
assert cap.quota_rules == []
assert cap.metadata == {}
def test_with_quota_rules(self):
rules = [QuotaRule(dimension="credits", per_operation=1.0)]
cap = ModuleCapability(
name="generate",
description="生成功能",
quota_rules=rules,
)
assert cap.description == "生成功能"
assert len(cap.quota_rules) == 1
assert cap.quota_rules[0].dimension == "credits"
def test_with_metadata(self):
cap = ModuleCapability(
name="export",
metadata={"format": "mp4", "max_resolution": "1080p"},
)
assert cap.metadata["format"] == "mp4"
assert cap.metadata["max_resolution"] == "1080p"
# ============================================================
# Module
# ============================================================
class TestModuleDefaults:
"""Module 数据类默认值"""
def test_required_name(self):
mod = Module(name="ai_voice")
assert mod.name == "ai_voice"
def test_default_version(self):
mod = Module(name="test")
assert mod.version == "1.0.0"
def test_default_description(self):
mod = Module(name="test")
assert mod.description == ""
def test_default_capabilities_empty(self):
mod = Module(name="test")
assert mod.capabilities == []
def test_default_dependencies_empty(self):
mod = Module(name="test")
assert mod.dependencies == []
def test_default_status_registered(self):
mod = Module(name="test")
assert mod.status == ModuleStatus.REGISTERED
def test_default_config_empty(self):
mod = Module(name="test")
assert mod.config == {}
def test_full_module(self):
cap = ModuleCapability(name="do_something")
mod = Module(
name="full_module",
version="2.0.0",
description="完整模块",
capabilities=[cap],
dependencies=["dep1", "dep2"],
status=ModuleStatus.ACTIVE,
config={"key": "value"},
)
assert mod.version == "2.0.0"
assert mod.description == "完整模块"
assert len(mod.capabilities) == 1
assert mod.dependencies == ["dep1", "dep2"]
assert mod.status == ModuleStatus.ACTIVE
assert mod.config["key"] == "value"
class TestModuleActivate:
"""Module.activate 状态转换"""
def test_activate_from_registered(self):
mod = Module(name="test")
mod.activate()
assert mod.status == ModuleStatus.ACTIVE
def test_activate_from_disabled(self):
mod = Module(name="test", status=ModuleStatus.DISABLED)
mod.activate()
assert mod.status == ModuleStatus.ACTIVE
def test_activate_from_error_stays_error(self):
mod = Module(name="test", status=ModuleStatus.ERROR)
mod.activate()
# error 状态不可激活
assert mod.status == ModuleStatus.ERROR
def test_activate_already_active(self):
mod = Module(name="test", status=ModuleStatus.ACTIVE)
mod.activate()
assert mod.status == ModuleStatus.ACTIVE
class TestModuleDisable:
"""Module.disable 状态转换"""
def test_disable_from_registered(self):
mod = Module(name="test")
mod.disable()
assert mod.status == ModuleStatus.DISABLED
def test_disable_from_active(self):
mod = Module(name="test", status=ModuleStatus.ACTIVE)
mod.disable()
assert mod.status == ModuleStatus.DISABLED
def test_disable_from_error(self):
mod = Module(name="test", status=ModuleStatus.ERROR)
mod.disable()
assert mod.status == ModuleStatus.DISABLED
def test_disable_already_disabled(self):
mod = Module(name="test", status=ModuleStatus.DISABLED)
mod.disable()
assert mod.status == ModuleStatus.DISABLED
# ============================================================
# ModuleRegistry - 基础操作
# ============================================================
class TestModuleRegistryBasic:
"""ModuleRegistry 基础操作"""
def test_empty_registry(self):
registry = ModuleRegistry()
assert registry.list_modules() == []
assert registry.get_active_capabilities() == {}
def test_register_single_module(self):
registry = ModuleRegistry()
mod = Module(name="test_mod")
registry.register(mod)
assert registry.get("test_mod") is mod
def test_register_duplicate_raises(self):
registry = ModuleRegistry()
registry.register(Module(name="test_mod"))
with pytest.raises(ValueError, match="already registered"):
registry.register(Module(name="test_mod"))
def test_get_nonexistent_returns_none(self):
registry = ModuleRegistry()
assert registry.get("no_such_module") is None
def test_unregister_success(self):
registry = ModuleRegistry()
registry.register(Module(name="test_mod"))
registry.unregister("test_mod")
assert registry.get("test_mod") is None
def test_unregister_nonexistent_raises(self):
registry = ModuleRegistry()
with pytest.raises(KeyError, match="not found"):
registry.unregister("no_such_module")
def test_unregister_with_dependents_raises(self):
registry = ModuleRegistry()
registry.register(Module(name="base_module"))
registry.register(Module(name="dependent_module", dependencies=["base_module"]))
with pytest.raises(ValueError, match="depended on by"):
registry.unregister("base_module")
def test_clear(self):
registry = ModuleRegistry()
registry.register(Module(name="mod1"))
registry.register(Module(name="mod2"))
registry.clear()
assert registry.list_modules() == []
# ============================================================
# ModuleRegistry - 自动激活 & 依赖
# ============================================================
class TestModuleRegistryAutoActivate:
"""注册时自动激活逻辑"""
def test_no_deps_auto_activates(self):
registry = ModuleRegistry()
mod = Module(name="standalone")
registry.register(mod)
assert mod.status == ModuleStatus.ACTIVE
def test_with_deps_all_satisfied_auto_activates(self):
registry = ModuleRegistry()
registry.register(Module(name="base")) # 无依赖,自动激活
dep_mod = Module(name="dependent", dependencies=["base"])
registry.register(dep_mod)
assert dep_mod.status == ModuleStatus.ACTIVE
def test_with_deps_not_satisfied_stays_registered(self):
registry = ModuleRegistry()
mod = Module(name="dependent", dependencies=["missing_dep"])
registry.register(mod)
# 依赖不满足,保持 REGISTERED
assert mod.status == ModuleStatus.REGISTERED
def test_later_dep_registered_manual_activate(self):
"""先注册依赖模块,再注册被依赖模块时不自动激活前者
(需要手动或在注册完所有模块后调用 check_dependencies + activate"""
registry = ModuleRegistry()
# 先注册依赖方(依赖未满足,不激活)
dependent = Module(name="dependent", dependencies=["base"])
registry.register(dependent)
assert dependent.status == ModuleStatus.REGISTERED
# 再注册被依赖方
base = Module(name="base")
registry.register(base)
assert base.status == ModuleStatus.ACTIVE
# 依赖方仍然是 REGISTERED(不会自动激活)
assert dependent.status == ModuleStatus.REGISTERED
class TestModuleRegistryCheckDependencies:
"""check_dependencies 依赖检查"""
def test_module_not_found_returns_false(self):
registry = ModuleRegistry()
assert registry.check_dependencies("nonexistent") is False
def test_no_deps_returns_true(self):
registry = ModuleRegistry()
registry.register(Module(name="standalone"))
assert registry.check_dependencies("standalone") is True
def test_all_deps_active_returns_true(self):
registry = ModuleRegistry()
registry.register(Module(name="dep1"))
registry.register(Module(name="dep2"))
registry.register(Module(name="main", dependencies=["dep1", "dep2"]))
# main 在注册时因依赖满足已自动激活
assert registry.check_dependencies("main") is True
def test_dep_not_registered_returns_false(self):
registry = ModuleRegistry()
mod = Module(name="main", dependencies=["missing"])
registry.register(mod)
assert registry.check_dependencies("main") is False
def test_dep_registered_but_not_active_returns_false(self):
registry = ModuleRegistry()
dep = Module(name="dep", status=ModuleStatus.DISABLED)
registry.register(dep)
# 手动设为 disabled(因为 register 时无依赖会自动激活)
dep.disable()
main = Module(name="main", dependencies=["dep"])
registry.register(main)
# 依赖未激活
assert registry.check_dependencies("main") is False
# ============================================================
# ModuleRegistry - list_modules & 状态过滤
# ============================================================
class TestModuleRegistryList:
"""list_modules 列表与过滤"""
def test_list_all(self):
registry = ModuleRegistry()
registry.register(Module(name="mod1"))
registry.register(Module(name="mod2"))
modules = registry.list_modules()
assert len(modules) == 2
names = {m.name for m in modules}
assert names == {"mod1", "mod2"}
def test_filter_by_active(self):
registry = ModuleRegistry()
registry.register(Module(name="active_mod")) # 自动激活
disabled = Module(name="disabled_mod")
registry.register(disabled)
disabled.disable()
active = registry.list_modules(status=ModuleStatus.ACTIVE)
assert len(active) == 1
assert active[0].name == "active_mod"
def test_filter_by_disabled(self):
registry = ModuleRegistry()
registry.register(Module(name="active_mod"))
disabled = Module(name="disabled_mod")
registry.register(disabled)
disabled.disable()
disabled_list = registry.list_modules(status=ModuleStatus.DISABLED)
assert len(disabled_list) == 1
assert disabled_list[0].name == "disabled_mod"
def test_filter_registered(self):
registry = ModuleRegistry()
# 有依赖未满足的模块保持 REGISTERED
mod = Module(name="waiting_mod", dependencies=["missing"])
registry.register(mod)
registered = registry.list_modules(status=ModuleStatus.REGISTERED)
assert len(registered) == 1
assert registered[0].name == "waiting_mod"
# ============================================================
# ModuleRegistry - 能力发现
# ============================================================
class TestModuleRegistryCapabilities:
"""能力发现:has_capability / get_capability / get_quota_rules"""
def test_has_capability_true(self):
registry = ModuleRegistry()
registry.register(
Module(
name="voice_module",
capabilities=[ModuleCapability(name="generate_voice")],
)
)
assert registry.has_capability("generate_voice") is True
def test_has_capability_false(self):
registry = ModuleRegistry()
registry.register(
Module(
name="voice_module",
capabilities=[ModuleCapability(name="generate_voice")],
)
)
assert registry.has_capability("generate_video") is False
def test_has_capability_inactive_module_not_counted(self):
registry = ModuleRegistry()
mod = Module(
name="inactive_mod",
capabilities=[ModuleCapability(name="secret_cap")],
)
registry.register(mod)
mod.disable()
assert registry.has_capability("secret_cap") is False
def test_get_capability_returns_first_match(self):
registry = ModuleRegistry()
cap1 = ModuleCapability(name="export", description="导出1")
cap2 = ModuleCapability(name="export", description="导出2")
registry.register(Module(name="mod1", capabilities=[cap1]))
registry.register(Module(name="mod2", capabilities=[cap2]))
result = registry.get_capability("export")
assert result is not None
assert result.name == "export"
# 返回第一个匹配的(mod1
assert result.description == "导出1"
def test_get_capability_nonexistent_returns_none(self):
registry = ModuleRegistry()
assert registry.get_capability("no_such_cap") is None
def test_get_quota_rules(self):
rules = [
QuotaRule(dimension="credits", per_operation=1.0),
QuotaRule(dimension="storage", per_operation=0.5),
]
registry = ModuleRegistry()
registry.register(
Module(
name="voice_mod",
capabilities=[ModuleCapability(name="gen", quota_rules=rules)],
)
)
result = registry.get_quota_rules("gen")
assert len(result) == 2
assert result[0].dimension == "credits"
assert result[1].dimension == "storage"
def test_get_quota_rules_nonexistent_returns_empty(self):
registry = ModuleRegistry()
assert registry.get_quota_rules("no_cap") == []
# ============================================================
# ModuleRegistry - get_active_capabilities
# ============================================================
class TestModuleRegistryActiveCapabilities:
"""get_active_capabilities 已激活能力汇总"""
def test_empty_registry(self):
registry = ModuleRegistry()
assert registry.get_active_capabilities() == {}
def test_single_module_with_caps(self):
registry = ModuleRegistry()
registry.register(
Module(
name="voice_mod",
capabilities=[
ModuleCapability(name="generate_voice"),
ModuleCapability(name="clone_voice"),
],
)
)
result = registry.get_active_capabilities()
assert "voice_mod" in result
assert set(result["voice_mod"]) == {"generate_voice", "clone_voice"}
def test_skips_inactive_modules(self):
registry = ModuleRegistry()
registry.register(
Module(
name="active_mod",
capabilities=[ModuleCapability(name="active_cap")],
)
)
inactive = Module(
name="inactive_mod",
capabilities=[ModuleCapability(name="inactive_cap")],
)
registry.register(inactive)
inactive.disable()
result = registry.get_active_capabilities()
assert "active_mod" in result
assert "inactive_mod" not in result
def test_skips_modules_without_caps(self):
registry = ModuleRegistry()
registry.register(Module(name="no_cap_mod"))
result = registry.get_active_capabilities()
assert "no_cap_mod" not in result
def test_multiple_modules(self):
registry = ModuleRegistry()
registry.register(
Module(
name="mod1",
capabilities=[ModuleCapability(name="cap_a")],
)
)
registry.register(
Module(
name="mod2",
capabilities=[ModuleCapability(name="cap_b"), ModuleCapability(name="cap_c")],
)
)
result = registry.get_active_capabilities()
assert len(result) == 2
assert result["mod1"] == ["cap_a"]
assert set(result["mod2"]) == {"cap_b", "cap_c"}
# ============================================================
# 全局单例
# ============================================================
class TestGlobalSingleton:
"""全局 module_registry 单例"""
def test_singleton_exists(self):
assert module_registry is not None
assert isinstance(module_registry, ModuleRegistry)
def test_singleton_is_same_instance(self):
from packages.infrastructure.module_registry import module_registry as mr2
assert module_registry is mr2
+337
View File
@@ -0,0 +1,337 @@
"""
pagination 通用分页器单元测试
覆盖:
- PaginationParams: 默认值/边界/校验/offset/limit
- PaginationMeta: from_params 各种边界场景
- PaginatedResponse: create 工厂方法
- paginate: 内存分页函数
"""
import pytest
from pydantic import ValidationError
from packages.application.common.pagination import (
PaginatedResponse,
PaginationMeta,
PaginationParams,
paginate,
)
# ============================================================
# PaginationParams
# ============================================================
class TestPaginationParamsDefaults:
"""默认值测试"""
def test_default_page_is_1(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()
assert params.offset == 0
def test_default_limit_is_20(self):
params = PaginationParams()
assert params.limit == 20
class TestPaginationParamsValidation:
"""参数校验"""
@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):
with pytest.raises(ValidationError):
PaginationParams(page=0)
def test_page_negative_raises(self):
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):
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):
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):
params = PaginationParams(page_size=100)
assert params.limit == 100
# ============================================================
# PaginationMeta.from_params
# ============================================================
class TestPaginationMetaFromParams:
"""from_params 工厂方法"""
def test_empty_total_zero(self):
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)
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
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)
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):
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
def test_first_page(self):
items = list(range(50))
params = PaginationParams(page=1, page_size=20)
result = paginate(items, params)
assert result.data == list(range(20))
assert result.pagination.total == 50
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)
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
def test_last_page(self):
items = list(range(50))
params = PaginationParams(page=3, page_size=20)
result = paginate(items, params)
assert result.data == list(range(40, 50))
assert len(result.data) == 10
assert result.pagination.has_next is False
assert result.pagination.has_prev is True
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)
result = paginate(items, params)
assert result.data == items
assert result.pagination.total_pages == 1
assert result.pagination.has_next is False
def test_single_item(self):
items = [42]
params = PaginationParams(page=1, page_size=20)
result = paginate(items, params)
assert result.data == [42]
assert result.pagination.total == 1
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
+539
View File
@@ -0,0 +1,539 @@
"""
SharedStorageService 单元测试
重点覆盖纯逻辑部分:
- _normalize_storage_key: URL提取 + URL解码
- _is_local_generated_url: 本地生成URL判断
- get_url: 公共URL拼接
- create_direct_upload_post: policy + HMAC签名
- get_download_url: bucket=None时的fallback
- 未配置OSS时的错误处理
- 单例模式
"""
import base64
import hashlib
import hmac
import json
import os
from unittest.mock import MagicMock, patch
import pytest
from packages.shared.storage import (
SharedStorageService,
get_shared_storage_service,
get_storage_service,
)
# ============================================================
# Fixtures
# ============================================================
def _make_service(
bucket_name="test-bucket",
endpoint="oss-cn-hangzhou.aliyuncs.com",
access_key_id="test-key-id",
access_key_secret="test-key-secret",
local_url_prefix="/generated-files",
with_bucket=True,
):
"""创建一个 SharedStorageService 实例,mock 掉 oss2 和 settings。"""
mock_settings = MagicMock()
mock_settings.oss_bucket_name = bucket_name
mock_settings.oss_endpoint = endpoint
mock_settings.oss_access_key_id = access_key_id
mock_settings.oss_access_key_secret = access_key_secret
mock_bucket = MagicMock() if with_bucket else None
with (
patch("packages.shared.storage.get_shared_settings", return_value=mock_settings),
patch.dict(os.environ, {"GENERATED_FILES_URL_PREFIX": local_url_prefix}, clear=False),
):
if with_bucket:
with patch("packages.shared.storage.oss2") as mock_oss2:
mock_oss2.Auth.return_value = MagicMock()
mock_oss2.Bucket.return_value = mock_bucket
service = SharedStorageService()
service.bucket = mock_bucket
return service, mock_bucket, mock_settings
else:
service = SharedStorageService()
service.bucket = None
return service, None, mock_settings
# ============================================================
# _normalize_storage_key
# ============================================================
class TestNormalizeStorageKey:
"""_normalize_storage_key URL 提取与解码"""
def test_plain_key_returns_as_is(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("videos/clip.mp4")
assert result == "videos/clip.mp4"
def test_key_with_leading_slash_stripped(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("/videos/clip.mp4")
assert result == "videos/clip.mp4"
def test_https_url_extracts_path(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("https://test-bucket.oss-cn-hangzhou.aliyuncs.com/videos/clip.mp4")
assert result == "videos/clip.mp4"
def test_http_url_extracts_path(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("http://test-bucket.oss-cn-hangzhou.aliyuncs.com/audio/voice.mp3")
assert result == "audio/voice.mp3"
def test_url_with_query_strips_query(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("https://bucket.oss-cn.com/file.mp4?signature=abc&expires=123")
assert result == "file.mp4"
def test_url_with_leading_slash_in_path(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("https://bucket.oss.com//double/slash.jpg")
assert result == "double/slash.jpg"
def test_url_decodes_percent_encoded_spaces(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("https://bucket.oss.com/my%20video.mp4")
assert result == "my video.mp4"
def test_url_decodes_percent_encoded_chinese(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("https://bucket.oss.com/%E4%B8%AD%E6%96%87.mp4")
assert result == "中文.mp4"
def test_url_with_special_chars_decoded(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("https://bucket.oss.com/file%281%29.jpg")
assert result == "file(1).jpg"
def test_plain_key_with_percent_not_decoded(self):
"""原始 key 不以 http 开头,不做 URL 解码,直接 lstrip('/')"""
service, _, _ = _make_service()
result = service._normalize_storage_key("file%20name.mp4")
# 不是 URL,直接返回(去掉前导/)
assert result == "file%20name.mp4"
def test_empty_string(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("")
assert result == ""
def test_root_slash_url(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("https://bucket.oss.com/")
assert result == ""
def test_nested_path_url(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("https://bucket.oss.com/a/b/c/d/file.txt")
assert result == "a/b/c/d/file.txt"
def test_url_with_port(self):
service, _, _ = _make_service()
result = service._normalize_storage_key("https://bucket.oss.com:443/file.txt")
assert result == "file.txt"
# ============================================================
# _is_local_generated_url
# ============================================================
class TestIsLocalGeneratedUrl:
"""_is_local_generated_url 本地URL判断"""
def test_local_prefix_returns_true(self):
service, _, _ = _make_service()
assert service._is_local_generated_url("/generated-files/abc.mp4") is True
def test_relative_local_returns_true(self):
service, _, _ = _make_service()
# 没有 scheme,直接用原字符串匹配
assert service._is_local_generated_url("/generated-files/out.mp4") is True
def test_full_url_with_local_path_returns_true(self):
service, _, _ = _make_service()
assert service._is_local_generated_url("https://example.com/generated-files/abc.mp4") is True
def test_other_path_returns_false(self):
service, _, _ = _make_service()
assert service._is_local_generated_url("/videos/abc.mp4") is False
def test_empty_string_returns_false(self):
service, _, _ = _make_service()
assert service._is_local_generated_url("") is False
def test_custom_prefix(self):
service, _, _ = _make_service(local_url_prefix="/custom-prefix")
assert service._is_local_generated_url("/custom-prefix/file.mp4") is True
assert service._is_local_generated_url("/generated-files/file.mp4") is False
# ============================================================
# get_url
# ============================================================
class TestGetUrl:
"""get_url 公共URL拼接"""
def test_returns_public_url_plus_key(self):
service, _, _ = _make_service()
result = service.get_url("videos/test.mp4")
assert result == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/videos/test.mp4"
def test_empty_key(self):
service, _, _ = _make_service()
result = service.get_url("")
assert result == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/"
def test_custom_bucket_and_endpoint(self):
service, _, _ = _make_service(
bucket_name="my-bucket",
endpoint="oss-us-east-1.aliyuncs.com",
)
result = service.get_url("file.txt")
assert result == "https://my-bucket.oss-us-east-1.aliyuncs.com/file.txt"
# ============================================================
# create_direct_upload_post
# ============================================================
class TestCreateDirectUploadPost:
"""create_direct_upload_post 直传表单生成"""
def test_returns_dict_with_expected_keys(self):
service, _, _ = _make_service()
result = service.create_direct_upload_post(
storage_key="uploads/test.jpg",
content_type="image/jpeg",
max_size_bytes=10 * 1024 * 1024,
expires_seconds=3600,
)
assert "url" in result
assert "method" in result
assert "storage_key" in result
assert "expires_at" in result
assert "fields" in result
def test_method_is_post(self):
service, _, _ = _make_service()
result = service.create_direct_upload_post("uploads/a.jpg", "image/jpeg", 1024, 3600)
assert result["method"] == "POST"
def test_url_is_public_url(self):
service, _, _ = _make_service()
result = service.create_direct_upload_post("uploads/a.jpg", "image/jpeg", 1024, 3600)
assert result["url"] == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com"
def test_storage_key_normalized(self):
service, _, _ = _make_service()
result = service.create_direct_upload_post("/uploads/test.jpg", "image/jpeg", 1024, 3600)
assert result["storage_key"] == "uploads/test.jpg"
assert result["fields"]["key"] == "uploads/test.jpg"
def test_fields_contain_required_keys(self):
service, _, _ = _make_service()
result = service.create_direct_upload_post("uploads/a.jpg", "image/jpeg", 1024, 3600)
fields = result["fields"]
assert fields["key"] == "uploads/a.jpg"
assert fields["OSSAccessKeyId"] == "test-key-id"
assert fields["success_action_status"] == "201"
assert fields["Content-Type"] == "image/jpeg"
assert "policy" in fields
assert "Signature" in fields
def test_policy_signature_is_valid_hmac_sha1(self):
"""验证 HMAC-SHA1 签名是否正确"""
secret = "my-secret-key-123"
service, _, _ = _make_service(access_key_secret=secret)
result = service.create_direct_upload_post("uploads/a.jpg", "image/jpeg", 1024, 3600)
policy = result["fields"]["policy"]
signature = result["fields"]["Signature"]
# 手动计算签名验证
expected = base64.b64encode(
hmac.new(secret.encode("utf-8"), policy.encode("utf-8"), hashlib.sha1).digest()
).decode("ascii")
assert signature == expected
def test_policy_contains_bucket_and_key(self):
service, _, _ = _make_service(bucket_name="my-bucket")
result = service.create_direct_upload_post("uploads/photo.png", "image/png", 2048, 1800)
policy = json.loads(base64.b64decode(result["fields"]["policy"]))
conditions = policy["conditions"]
assert {"bucket": "my-bucket"} in conditions
assert {"key": "uploads/photo.png"} in conditions
def test_policy_contains_content_length_range(self):
service, _, _ = _make_service()
result = service.create_direct_upload_post("uploads/a.jpg", "image/jpeg", 5242880, 3600)
policy = json.loads(base64.b64decode(result["fields"]["policy"]))
conditions = policy["conditions"]
size_condition = [c for c in conditions if isinstance(c, list) and c[0] == "content-length-range"]
assert len(size_condition) == 1
assert size_condition[0][1] == 1
assert size_condition[0][2] == 5242880
def test_policy_content_type_starts_with(self):
service, _, _ = _make_service()
result = service.create_direct_upload_post("uploads/a.jpg", "image/jpeg", 1024, 3600)
policy = json.loads(base64.b64decode(result["fields"]["policy"]))
conditions = policy["conditions"]
ct_condition = [c for c in conditions if isinstance(c, list) and c[0] == "starts-with"]
assert len(ct_condition) == 1
assert ct_condition[0][1] == "$Content-Type"
assert ct_condition[0][2] == "image/"
def test_policy_has_expiration(self):
service, _, _ = _make_service()
result = service.create_direct_upload_post("uploads/a.jpg", "image/jpeg", 1024, 3600)
policy = json.loads(base64.b64decode(result["fields"]["policy"]))
assert "expiration" in policy
# ISO 8601 格式
assert policy["expiration"].endswith("Z")
def test_non_uploads_key_raises_value_error(self):
service, _, _ = _make_service()
with pytest.raises(ValueError, match="uploads/"):
service.create_direct_upload_post("videos/a.mp4", "video/mp4", 1024, 3600)
def test_no_credentials_raises_runtime_error(self):
service, _, _ = _make_service(access_key_id="", access_key_secret="", with_bucket=False)
with pytest.raises(RuntimeError, match="not configured"):
service.create_direct_upload_post("uploads/a.jpg", "image/jpeg", 1024, 3600)
def test_url_normalized_key_in_uploads(self):
service, _, _ = _make_service()
# URL 形式的 key 被 normalize 后如果在 uploads/ 下应该可以
result = service.create_direct_upload_post(
"https://test-bucket.oss-cn-hangzhou.aliyuncs.com/uploads/from_url.jpg",
"image/jpeg",
1024,
3600,
)
assert result["storage_key"] == "uploads/from_url.jpg"
# ============================================================
# get_download_url (bucket=None 时的 fallback)
# ============================================================
class TestGetDownloadUrlFallback:
"""get_download_url 在 bucket 未配置时的 fallback 逻辑"""
def test_no_bucket_local_url_returns_as_is(self):
service, _, _ = _make_service(with_bucket=False)
result = service.get_download_url("/generated-files/test.mp4")
assert result == "/generated-files/test.mp4"
def test_no_bucket_regular_key_returns_public_url(self):
service, _, _ = _make_service(with_bucket=False)
result = service.get_download_url("videos/test.mp4")
assert result == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/videos/test.mp4"
def test_no_bucket_url_input_normalized(self):
service, _, _ = _make_service(with_bucket=False)
result = service.get_download_url("https://test-bucket.oss-cn-hangzhou.aliyuncs.com/videos/test.mp4")
assert result == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/videos/test.mp4"
def test_with_bucket_calls_sign_url(self):
service, mock_bucket, _ = _make_service(with_bucket=True)
mock_bucket.sign_url.return_value = "https://signed-url.com/file?sig=abc"
result = service.get_download_url("videos/test.mp4", expires_seconds=7200)
mock_bucket.sign_url.assert_called_once_with("GET", "videos/test.mp4", 7200)
assert result == "https://signed-url.com/file?sig=abc"
def test_sign_url_exception_falls_back_to_public_url(self):
service, mock_bucket, _ = _make_service(with_bucket=True)
mock_bucket.sign_url.side_effect = Exception("sign error")
result = service.get_download_url("videos/test.mp4")
assert result == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/videos/test.mp4"
# ============================================================
# 未配置 OSS 时的错误处理
# ============================================================
class TestNoBucketErrorHandling:
"""bucket=None 时的错误处理"""
def test_upload_file_raises(self):
service, _, _ = _make_service(with_bucket=False)
with pytest.raises(RuntimeError, match="not configured"):
service.upload_file("/tmp/test.txt", "uploads/test.txt")
def test_download_file_raises(self):
service, _, _ = _make_service(with_bucket=False)
with pytest.raises(RuntimeError, match="not configured"):
service.download_file("uploads/test.txt", "/tmp/test.txt")
def test_delete_file_silent_noop(self):
service, _, _ = _make_service(with_bucket=False)
# 不抛异常
result = service.delete_file("uploads/test.txt")
assert result is None
def test_file_exists_returns_false(self):
service, _, _ = _make_service(with_bucket=False)
assert service.file_exists("uploads/test.txt") is False
# ============================================================
# upload_file / delete_file / file_exists 正常路径
# ============================================================
class TestBucketOperations:
"""有 bucket 时的操作调用验证"""
def test_upload_file_with_path_string(self):
service, mock_bucket, _ = _make_service()
result = service.upload_file("/tmp/file.txt", "uploads/file.txt", "text/plain")
mock_bucket.put_object_from_file.assert_called_once()
args = mock_bucket.put_object_from_file.call_args
assert args[0][0] == "uploads/file.txt"
assert args[0][1] == "/tmp/file.txt"
assert result.startswith("https://test-bucket.")
def test_upload_file_with_file_object(self):
service, mock_bucket, _ = _make_service()
mock_file = MagicMock()
result = service.upload_file(mock_file, "uploads/file.bin", "application/octet-stream")
mock_file.seek.assert_called_once_with(0)
mock_bucket.put_object.assert_called_once()
assert result.startswith("https://test-bucket.")
def test_delete_file_calls_bucket(self):
service, mock_bucket, _ = _make_service()
service.delete_file("uploads/test.txt")
mock_bucket.delete_object.assert_called_once_with("uploads/test.txt")
def test_delete_file_exception_logged_not_raised(self):
service, mock_bucket, _ = _make_service()
mock_bucket.delete_object.side_effect = Exception("delete error")
# 不抛异常
service.delete_file("uploads/test.txt")
def test_file_exists_delegates_to_bucket(self):
service, mock_bucket, _ = _make_service()
mock_bucket.object_exists.return_value = True
assert service.file_exists("some/key") is True
mock_bucket.object_exists.assert_called_once_with("some/key")
def test_file_exists_false(self):
service, mock_bucket, _ = _make_service()
mock_bucket.object_exists.return_value = False
assert service.file_exists("some/key") is False
# ============================================================
# 单例 & 兼容别名
# ============================================================
class TestSingleton:
"""get_shared_storage_service 单例模式"""
def test_get_storage_service_is_alias(self):
# 两个函数返回同一个实例
with patch("packages.shared.storage._storage_service", None):
with patch("packages.shared.storage.SharedStorageService") as mock_cls:
mock_instance = MagicMock()
mock_cls.return_value = mock_instance
svc1 = get_shared_storage_service()
svc2 = get_storage_service()
assert svc1 is svc2
# 因为是同一个单例,类只实例化一次
assert mock_cls.call_count == 1
# ============================================================
# __init__ endpoint 处理
# ============================================================
class TestInitEndpointHandling:
"""初始化时 endpoint https 前缀处理"""
def test_endpoint_without_https_gets_prefix(self):
mock_settings = MagicMock()
mock_settings.oss_bucket_name = "test-bucket"
mock_settings.oss_endpoint = "oss-cn-hangzhou.aliyuncs.com"
mock_settings.oss_access_key_id = "key-id"
mock_settings.oss_access_key_secret = "key-secret"
with (
patch("packages.shared.storage.get_shared_settings", return_value=mock_settings),
patch("packages.shared.storage.oss2") as mock_oss2,
):
mock_oss2.Bucket.return_value = MagicMock()
service = SharedStorageService()
# 验证 Bucket 构造时 endpoint 带了 https://
call_args = mock_oss2.Bucket.call_args
assert call_args[0][1] == "https://oss-cn-hangzhou.aliyuncs.com"
def test_endpoint_with_https_keeps_as_is(self):
mock_settings = MagicMock()
mock_settings.oss_bucket_name = "test-bucket"
mock_settings.oss_endpoint = "https://oss-cn-hangzhou.aliyuncs.com"
mock_settings.oss_access_key_id = "key-id"
mock_settings.oss_access_key_secret = "key-secret"
with (
patch("packages.shared.storage.get_shared_settings", return_value=mock_settings),
patch("packages.shared.storage.oss2") as mock_oss2,
):
mock_oss2.Bucket.return_value = MagicMock()
service = SharedStorageService()
call_args = mock_oss2.Bucket.call_args
assert call_args[0][1] == "https://oss-cn-hangzhou.aliyuncs.com"
def test_endpoint_with_http_keeps_as_is(self):
mock_settings = MagicMock()
mock_settings.oss_bucket_name = "test-bucket"
mock_settings.oss_endpoint = "http://oss-cn-hangzhou.aliyuncs.com"
mock_settings.oss_access_key_id = "key-id"
mock_settings.oss_access_key_secret = "key-secret"
with (
patch("packages.shared.storage.get_shared_settings", return_value=mock_settings),
patch("packages.shared.storage.oss2") as mock_oss2,
):
mock_oss2.Bucket.return_value = MagicMock()
service = SharedStorageService()
call_args = mock_oss2.Bucket.call_args
assert call_args[0][1] == "http://oss-cn-hangzhou.aliyuncs.com"
+315
View File
@@ -0,0 +1,315 @@
"""
text_splitter 长文本分段工具单元测试
覆盖:
- 空文本 / 短文本
- 句子边界分段(。!?;\n . ! ? ;
- 超长句子硬切
- 过短段落合并
- max_chars 参数
- 中英文混合
"""
import pytest
from packages.application.tts_job.text_splitter import split_text
# ============================================================
# 基础场景
# ============================================================
class TestBasicCases:
"""基础场景"""
def test_empty_text_returns_empty_list(self):
assert split_text("") == []
def test_whitespace_only_returns_empty(self):
assert split_text(" \n\n ") == []
def test_short_text_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)
assert len(result) == 1
assert len(result[0]) == 500
def test_text_stripped(self):
text = " 你好世界。 "
result = split_text(text, max_chars=500)
assert result == ["你好世界。"]
# ============================================================
# 句子边界分段
# ============================================================
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_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
# ============================================================
# 边界情况
# ============================================================
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
for seg in result:
assert len(seg) <= 200
+553 -252
View File
@@ -1,296 +1,597 @@
"""URL 安全校验工具单元测试 — SSRF 防护."""
"""
url_security URL安全校验单元测试
from __future__ import annotations
覆盖:
- validate_url_safety: scheme/主机/端口/SSRF/内网域名/白名单
- is_url_safe: 便捷函数
- UrlSecurityError / NoRedirectHandler
- _validate_magic_number: 文件魔数校验
- safe_download_file / safe_download_bytes: mock 网络测试
"""
import os
import sys
import unittest
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker"))
import shutil
import tempfile
from unittest.mock import MagicMock, patch
from video_processing.url_security import ( # noqa: E402
import pytest
from packages.shared.url_security import (
ALLOWED_AUDIO_MIME_TYPES,
ALLOWED_IMAGE_MIME_TYPES,
ALLOWED_PORTS,
ALLOWED_SCHEMES,
MAX_URL_LENGTH,
NoRedirectHandler,
UrlSecurityError,
_check_internal_hostnames,
_is_trusted_domain,
_validate_magic_number,
is_url_safe,
safe_download_bytes,
safe_download_file,
validate_url_safety,
)
class TestUrlSecurityValidation(unittest.TestCase):
"""URL 安全校验测试."""
# ── Scheme 白名单 ──────────────────────────────────────────────────────
def test_http_scheme_allowed(self):
"""HTTP scheme 应该被允许."""
result = validate_url_safety("http://example.com/test", purpose="test")
self.assertEqual(result, "http://example.com/test")
def test_https_scheme_allowed(self):
"""HTTPS scheme 应该被允许."""
result = validate_url_safety("https://example.com/test", purpose="test")
self.assertEqual(result, "https://example.com/test")
def test_file_scheme_rejected(self):
"""file:// scheme 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("file:///etc/passwd", purpose="test")
def test_ftp_scheme_rejected(self):
"""ftp:// scheme 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("ftp://example.com/test", purpose="test")
def test_empty_scheme_rejected(self):
"""空 scheme 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("example.com/test", purpose="test")
# ── 端口白名单 ────────────────────────────────────────────────────────
def test_port_80_allowed(self):
"""端口 80 应该被允许."""
# 80端口是默认HTTP端口,不显式指定也可以
result = validate_url_safety("http://example.com:80/test", purpose="test")
self.assertIn("example.com", result)
def test_port_443_allowed(self):
"""端口 443 应该被允许."""
result = validate_url_safety("https://example.com:443/test", purpose="test")
self.assertIn("example.com", result)
def test_port_8080_rejected(self):
"""非标准端口 8080 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://example.com:8080/test", purpose="test")
def test_port_22_rejected(self):
"""SSH 端口 22 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://example.com:22/test", purpose="test")
# ── SSRF: 直接 IP 访问 ───────────────────────────────────────────────
def test_loopback_ip_rejected(self):
"""回环地址 127.0.0.1 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://127.0.0.1/test", purpose="test")
def test_private_ip_192_rejected(self):
"""内网地址 192.168.x.x 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://192.168.1.1/test", purpose="test")
def test_private_ip_10_rejected(self):
"""内网地址 10.x.x.x 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://10.0.0.1/test", purpose="test")
def test_private_ip_172_rejected(self):
"""内网地址 172.16.x.x 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://172.16.0.1/test", purpose="test")
def test_unspecified_ip_rejected(self):
"""未指定地址 0.0.0.0 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://0.0.0.0/test", purpose="test")
def test_ipv6_loopback_rejected(self):
"""IPv6 回环 ::1 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://[::1]/test", purpose="test")
def test_ipv6_link_local_rejected(self):
"""IPv6 链路本地地址应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://[fe80::1]/test", purpose="test")
# ── SSRF: 内网主机名 ─────────────────────────────────────────────────
def test_localhost_rejected(self):
"""localhost 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://localhost/test", purpose="test")
def test_local_domain_rejected(self):
""".local 域名应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://printer.local/test", purpose="test")
def test_internal_domain_rejected(self):
""".internal 域名应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http://db.internal/test", purpose="test")
# ── URL 格式校验 ─────────────────────────────────────────────────────
def test_empty_url_rejected(self):
"""空 URL 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("", purpose="test")
def test_none_url_rejected(self):
"""None URL 应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety(None, purpose="test") # type: ignore
def test_url_too_long_rejected(self):
"""超长 URL 应该被拒绝."""
long_url = "https://example.com/" + "a" * 3000
with self.assertRaises(UrlSecurityError):
validate_url_safety(long_url, purpose="test")
def test_no_hostname_rejected(self):
"""缺少主机名应该被拒绝."""
with self.assertRaises(UrlSecurityError):
validate_url_safety("http:///test", purpose="test")
# ── is_url_safe 便捷函数 ─────────────────────────────────────────────
def test_is_url_safe_true(self):
"""安全 URL 应该返回 True."""
self.assertTrue(is_url_safe("https://example.com/test", purpose="test"))
def test_is_url_safe_false(self):
"""不安全 URL 应该返回 False."""
self.assertFalse(is_url_safe("http://127.0.0.1/test", purpose="test"))
def test_is_url_safe_empty(self):
"""空 URL 应该返回 False."""
self.assertFalse(is_url_safe("", purpose="test"))
# ── validate_url_safety 基础校验 ─────────────────────────────────────────────
if __name__ == "__main__":
unittest.main()
class TestValidateUrlSafetyBasics:
"""URL 安全校验基础测试"""
def test_valid_http_url(self):
url = "http://example.com/file.mp4"
result = validate_url_safety(url)
assert result == url
def test_valid_https_url(self):
url = "https://example.com/file.mp4"
result = validate_url_safety(url)
assert result == url
def test_empty_url_raises(self):
with pytest.raises(UrlSecurityError, match="为空"):
validate_url_safety("")
def test_none_url_raises(self):
with pytest.raises(UrlSecurityError):
validate_url_safety(None)
def test_url_too_long_raises(self):
long_url = "https://example.com/" + "a" * 2050
with pytest.raises(UrlSecurityError, match="过长"):
validate_url_safety(long_url)
def test_url_at_max_length_ok(self):
base = "https://example.com/"
pad = "a" * (MAX_URL_LENGTH - len(base))
url = base + pad
assert len(url) <= MAX_URL_LENGTH
result = validate_url_safety(url)
assert result == url
def test_invalid_scheme_ftp_raises(self):
with pytest.raises(UrlSecurityError, match="scheme"):
validate_url_safety("ftp://example.com/file")
def test_invalid_scheme_file_raises(self):
with pytest.raises(UrlSecurityError, match="scheme"):
validate_url_safety("file:///etc/passwd")
def test_invalid_scheme_data_raises(self):
with pytest.raises(UrlSecurityError, match="scheme"):
validate_url_safety("data:text/html,<script>")
def test_missing_scheme_raises(self):
with pytest.raises(UrlSecurityError, match="scheme"):
validate_url_safety("example.com/file")
def test_missing_hostname_raises(self):
with pytest.raises(UrlSecurityError, match="主机名"):
validate_url_safety("http:///path")
def test_uppercase_scheme_normalized(self):
"""HTTP/HTTPS 大写也能通过"""
url = "HTTPS://example.com/file"
# scheme 检查用 lower 比较
result = validate_url_safety(url)
assert result == url
def test_default_port_80_ok(self):
url = "http://example.com:80/file"
result = validate_url_safety(url)
assert result == url
def test_default_port_443_ok(self):
url = "https://example.com:443/file"
result = validate_url_safety(url)
assert result == url
def test_non_standard_port_raises(self):
with pytest.raises(UrlSecurityError, match="端口"):
validate_url_safety("http://example.com:8080/file")
def test_port_22_ssh_raises(self):
with pytest.raises(UrlSecurityError, match="端口"):
validate_url_safety("http://example.com:22/file")
def test_port_3306_mysql_raises(self):
with pytest.raises(UrlSecurityError, match="端口"):
validate_url_safety("http://example.com:3306/file")
class TestSafeDownload(unittest.TestCase):
"""安全下载函数测试."""
# ── 内网主机名 / SSRF 防护 ───────────────────────────────────────────────────
def setUp(self):
self.temp_dir = tempfile.mkdtemp()
def tearDown(self):
shutil.rmtree(self.temp_dir, ignore_errors=True)
class TestInternalHostnameProtection:
"""内网主机名防护测试"""
def test_safe_download_file_rejects_ssrf(self):
"""SSRF 风险 URL 应该被拒绝下载."""
dest = os.path.join(self.temp_dir, "test.bin")
with self.assertRaises(UrlSecurityError):
safe_download_file("http://127.0.0.1/test", dest, purpose="test")
def test_localhost_raises(self):
with pytest.raises(UrlSecurityError, match="内部主机名"):
validate_url_safety("http://localhost/file")
def test_safe_download_bytes_rejects_ssrf(self):
"""SSRF 风险 URL 应该被拒绝下载(bytes 版本)."""
with self.assertRaises(UrlSecurityError):
safe_download_bytes("http://localhost/test", purpose="test")
def test_localhost_mixed_case_raises(self):
with pytest.raises(UrlSecurityError):
validate_url_safety("http://LocalHost/file")
def test_safe_download_file_size_limit(self):
"""超过大小限制应该被拒绝."""
# 用 mock server 测试太大的 content-length
dest = os.path.join(self.temp_dir, "test.bin")
# 直接验证参数:max_size=0 时任何下载都应超限
# (这里用一个可访问的 URL 并设置极小的限制)
# 为避免依赖外部网络,这里只测试函数参数传递
with unittest.mock.patch("urllib.request.build_opener") as mock_opener:
mock_resp = unittest.mock.MagicMock()
mock_resp.headers = {"Content-Length": "1000"}
mock_resp.read.return_value = b""
mock_opener.return_value.open.return_value = mock_resp
# 设置 max_size=500content-length=1000 应被拒绝
with self.assertRaises(UrlSecurityError):
safe_download_file(
"https://example.com/test",
def test_localhost_localdomain_raises(self):
with pytest.raises(UrlSecurityError):
_check_internal_hostnames("localhost.localdomain")
def test_metadata_hostname_raises(self):
with pytest.raises(UrlSecurityError):
_check_internal_hostnames("metadata")
def test_metadata_google_internal_raises(self):
with pytest.raises(UrlSecurityError):
_check_internal_hostnames("metadata.google.internal")
def test_dot_local_domain_raises(self):
with pytest.raises(UrlSecurityError, match="内网域名"):
validate_url_safety("http://myservice.local/file")
def test_dot_internal_domain_raises(self):
with pytest.raises(UrlSecurityError, match="内网域名"):
validate_url_safety("http://myservice.internal/file")
def test_dot_localdomain_raises(self):
with pytest.raises(UrlSecurityError):
_check_internal_hostnames("server.localdomain")
def test_loopback_ip_127_0_0_1_raises(self):
with pytest.raises(UrlSecurityError, match="直接 IP|回环"):
validate_url_safety("http://127.0.0.1/file")
def test_metadata_ip_169_254_raises(self):
"""云元数据服务 IP"""
with pytest.raises(UrlSecurityError):
validate_url_safety("http://169.254.169.254/latest/meta-data/")
def test_private_ip_10_raises(self):
with pytest.raises(UrlSecurityError):
validate_url_safety("http://10.0.0.1/file")
def test_private_ip_172_16_raises(self):
with pytest.raises(UrlSecurityError):
validate_url_safety("http://172.16.0.1/file")
def test_private_ip_192_168_raises(self):
with pytest.raises(UrlSecurityError):
validate_url_safety("http://192.168.1.1/file")
def test_unspecified_ip_0_0_0_0_raises(self):
with pytest.raises(UrlSecurityError):
validate_url_safety("http://0.0.0.0/file")
def test_ipv6_loopback_raises(self):
with pytest.raises(UrlSecurityError):
validate_url_safety("http://[::1]/file")
def test_public_ip_ok(self):
"""公网IP在ALLOW_DIRECT_IP默认关闭时应被拦截"""
# 默认 ALLOW_DIRECT_IP = false
with pytest.raises(UrlSecurityError, match="直接 IP"):
validate_url_safety("http://8.8.8.8/file")
# ── 可信域名白名单 ───────────────────────────────────────────────────────────
class TestTrustedDomains:
"""可信域名白名单测试"""
def test_is_trusted_domain_exact_match(self):
with patch("packages.shared.url_security.TRUSTED_DOMAINS", {"example.com", "cdn.example.org"}):
# 重新加载模块以应用环境变量不太现实,直接测函数
# 直接改全局状态再还原
import packages.shared.url_security as mod
from packages.shared.url_security import _is_trusted_domain
original = mod.TRUSTED_DOMAINS
mod.TRUSTED_DOMAINS = {"example.com", "cdn.example.org"}
try:
assert _is_trusted_domain("example.com") is True
assert _is_trusted_domain("cdn.example.org") is True
finally:
mod.TRUSTED_DOMAINS = original
def test_is_trusted_domain_subdomain(self):
import packages.shared.url_security as mod
original = mod.TRUSTED_DOMAINS
mod.TRUSTED_DOMAINS = {"example.com"}
try:
assert mod._is_trusted_domain("sub.example.com") is True
assert mod._is_trusted_domain("a.b.example.com") is True
finally:
mod.TRUSTED_DOMAINS = original
def test_is_trusted_domain_no_match(self):
import packages.shared.url_security as mod
original = mod.TRUSTED_DOMAINS
mod.TRUSTED_DOMAINS = {"example.com"}
try:
assert mod._is_trusted_domain("other.com") is False
assert mod._is_trusted_domain("notexample.com") is False
finally:
mod.TRUSTED_DOMAINS = original
def test_validate_with_trusted_domains_restricted(self):
"""白名单非空时,不在白名单中的域名被拒"""
import packages.shared.url_security as mod
original = mod.TRUSTED_DOMAINS
mod.TRUSTED_DOMAINS = {"trusted.com"}
try:
# 不在白名单中 - 在 _is_trusted_domain 检查时就被拒,不走 DNS
with pytest.raises(UrlSecurityError, match="白名单"):
validate_url_safety("https://untrusted.com/file")
# 在白名单中 - 需要 mock DNS 解析避免实际网络请求
with patch("packages.shared.url_security._check_ssrf_domain"):
result = validate_url_safety("https://trusted.com/file")
assert result == "https://trusted.com/file"
# 子域名
result = validate_url_safety("https://sub.trusted.com/file")
assert result == "https://sub.trusted.com/file"
finally:
mod.TRUSTED_DOMAINS = original
# ── is_url_safe 便捷函数 ─────────────────────────────────────────────────────
class TestIsUrlSafe:
"""is_url_safe 便捷函数测试"""
def test_safe_url_returns_true(self):
assert is_url_safe("https://example.com/file") is True
def test_unsafe_url_returns_false(self):
assert is_url_safe("http://localhost/file") is False
def test_empty_url_returns_false(self):
assert is_url_safe("") is False
def test_invalid_scheme_returns_false(self):
assert is_url_safe("ftp://example.com/file") is False
# ── UrlSecurityError 异常类 ───────────────────────────────────────────────────
class TestUrlSecurityError:
"""UrlSecurityError 异常类测试"""
def test_is_value_error_subclass(self):
assert issubclass(UrlSecurityError, ValueError)
def test_error_message(self):
err = UrlSecurityError("test message")
assert str(err) == "test message"
# ── NoRedirectHandler ────────────────────────────────────────────────────────
class TestNoRedirectHandler:
"""NoRedirectHandler 测试"""
def test_redirect_request_returns_none(self):
handler = NoRedirectHandler()
result = handler.redirect_request(
MagicMock(),
MagicMock(),
302,
"Found",
{"Location": "http://other.com"},
"http://other.com",
)
assert result is None
# ── 魔数校验 ─────────────────────────────────────────────────────────────────
class TestMagicNumberValidation:
"""文件魔数校验测试"""
def test_valid_png(self, tmp_path):
f = tmp_path / "test.png"
# PNG 文件头: 89 50 4E 47 0D 0A 1A 0A
f.write_bytes(b"\x89PNG\r\n\x1a\n" + b"\x00" * 100)
# 不抛异常 = 通过
_validate_magic_number(str(f), ALLOWED_IMAGE_MIME_TYPES)
def test_valid_jpeg(self, tmp_path):
f = tmp_path / "test.jpg"
# JPEG 文件头: FF D8 FF
f.write_bytes(b"\xff\xd8\xff\xe0" + b"\x00" * 100)
_validate_magic_number(str(f), ALLOWED_IMAGE_MIME_TYPES)
def test_valid_gif87a(self, tmp_path):
f = tmp_path / "test.gif"
f.write_bytes(b"GIF87a" + b"\x00" * 100)
_validate_magic_number(str(f), ALLOWED_IMAGE_MIME_TYPES)
def test_valid_gif89a(self, tmp_path):
f = tmp_path / "test.gif"
f.write_bytes(b"GIF89a" + b"\x00" * 100)
_validate_magic_number(str(f), ALLOWED_IMAGE_MIME_TYPES)
def test_valid_webp(self, tmp_path):
f = tmp_path / "test.webp"
# RIFF....WEBP
data = bytearray(b"RIFF")
data += b"\x00\x00\x00\x00" # size placeholder
data += b"WEBP"
data += b"\x00" * 100
f.write_bytes(bytes(data))
_validate_magic_number(str(f), ALLOWED_IMAGE_MIME_TYPES)
def test_valid_bmp(self, tmp_path):
f = tmp_path / "test.bmp"
f.write_bytes(b"BM" + b"\x00" * 100)
_validate_magic_number(str(f), ALLOWED_IMAGE_MIME_TYPES)
def test_valid_wav(self, tmp_path):
f = tmp_path / "test.wav"
# RIFF....WAVE
data = bytearray(b"RIFF")
data += b"\x00\x00\x00\x00"
data += b"WAVE"
data += b"\x00" * 100
f.write_bytes(bytes(data))
_validate_magic_number(str(f), ALLOWED_AUDIO_MIME_TYPES)
def test_valid_mp3_id3(self, tmp_path):
f = tmp_path / "test.mp3"
f.write_bytes(b"ID3\x03\x00\x00\x00\x00\x00\x00" + b"\x00" * 100)
_validate_magic_number(str(f), ALLOWED_AUDIO_MIME_TYPES)
def test_valid_mp3_adts(self, tmp_path):
f = tmp_path / "test.mp3"
f.write_bytes(b"\xff\xfb\x90\x00" + b"\x00" * 100)
_validate_magic_number(str(f), ALLOWED_AUDIO_MIME_TYPES)
def test_valid_ogg(self, tmp_path):
f = tmp_path / "test.ogg"
f.write_bytes(b"OggS\x00\x02\x00\x00" + b"\x00" * 100)
_validate_magic_number(str(f), ALLOWED_AUDIO_MIME_TYPES)
def test_valid_flac(self, tmp_path):
f = tmp_path / "test.flac"
f.write_bytes(b"fLaC" + b"\x00" * 100)
_validate_magic_number(str(f), ALLOWED_AUDIO_MIME_TYPES)
def test_invalid_file_content_raises(self, tmp_path):
f = tmp_path / "test.bin"
f.write_bytes(b"this is not an image file at all")
with pytest.raises(UrlSecurityError, match="魔数"):
_validate_magic_number(str(f), ALLOWED_IMAGE_MIME_TYPES)
def test_empty_file_raises(self, tmp_path):
f = tmp_path / "empty.bin"
f.write_bytes(b"")
with pytest.raises(UrlSecurityError, match="为空"):
_validate_magic_number(str(f), ALLOWED_IMAGE_MIME_TYPES)
def test_nonexistent_file_raises(self, tmp_path):
with pytest.raises(UrlSecurityError, match="读取文件头失败"):
_validate_magic_number(str(tmp_path / "no_such_file"), ALLOWED_IMAGE_MIME_TYPES)
def test_no_allowed_mime_types_skips(self, tmp_path):
"""allowed_mime_types 为空时跳过校验"""
f = tmp_path / "test.bin"
f.write_bytes(b"random data here")
# 不抛异常
_validate_magic_number(str(f), set())
def test_unknown_mime_types_skips(self, tmp_path):
"""没有已知魔数的 MIME 类型跳过校验"""
f = tmp_path / "test.bin"
f.write_bytes(b"random data")
_validate_magic_number(str(f), {"application/x-unknown-type"})
# ── safe_download_file (mock 网络) ───────────────────────────────────────────
class TestSafeDownloadFile:
"""safe_download_file 下载测试(mock 网络)"""
def test_download_success(self, tmp_path):
test_content = b"Hello, this is test file content!"
dest = str(tmp_path / "output.bin")
mock_resp = MagicMock()
mock_resp.headers = {"Content-Type": "application/octet-stream"}
mock_resp.read.side_effect = [test_content, b""]
with patch("packages.shared.url_security.NoRedirectHandler") as mock_handler_cls:
mock_handler = MagicMock()
mock_handler_cls.return_value = mock_handler
mock_opener = MagicMock()
mock_opener.open.return_value = mock_resp
with patch("urllib.request.build_opener", return_value=mock_opener):
size = safe_download_file(
"https://example.com/test.bin",
dest,
purpose="test",
max_size=500,
)
def test_safe_download_file_mime_rejected(self):
"""不允许的 MIME 类型应该被拒绝."""
dest = os.path.join(self.temp_dir, "test.bin")
with unittest.mock.patch("urllib.request.build_opener") as mock_opener:
mock_resp = unittest.mock.MagicMock()
mock_resp.headers = {"Content-Type": "text/html"}
mock_resp.read.return_value = b""
mock_opener.return_value.open.return_value = mock_resp
with self.assertRaises(UrlSecurityError):
safe_download_file(
"https://example.com/test.mp3",
dest,
purpose="test",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
)
assert size == len(test_content)
with open(dest, "rb") as f:
assert f.read() == test_content
def test_download_with_mime_check_passes(self, tmp_path):
# PNG 文件
test_content = b"\x89PNG\r\n\x1a\n" + b"\x00" * 200
dest = str(tmp_path / "test.png")
mock_resp = MagicMock()
mock_resp.headers = {"Content-Type": "image/png"}
mock_resp.read.side_effect = [test_content, b""]
with patch("urllib.request.build_opener") as mock_build:
mock_opener = MagicMock()
mock_opener.open.return_value = mock_resp
mock_build.return_value = mock_opener
def test_safe_download_file_mime_allowed(self):
"""允许的 MIME 类型应该通过."""
dest = os.path.join(self.temp_dir, "test.mp3")
with unittest.mock.patch("urllib.request.build_opener") as mock_opener:
mock_resp = unittest.mock.MagicMock()
mock_resp.headers = {"Content-Type": "audio/mpeg"}
mock_resp.read.side_effect = [b"ID3audio_data", b""]
mock_resp.geturl.return_value = "https://example.com/test.mp3"
mock_opener.return_value.open.return_value = mock_resp
size = safe_download_file(
"https://example.com/test.mp3",
"https://example.com/test.png",
dest,
purpose="test",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
allowed_mime_types={"image/png", "image/jpeg"},
)
self.assertEqual(size, 13)
self.assertTrue(os.path.exists(dest))
def test_safe_download_file_stream_size_limit(self):
"""流式下载时超过大小限制应该中断."""
dest = os.path.join(self.temp_dir, "test.bin")
with unittest.mock.patch("urllib.request.build_opener") as mock_opener:
mock_resp = unittest.mock.MagicMock()
mock_resp.headers = {}
# 每次返回 100 字节,max_size=500,第 6 次读取就超限
mock_resp.read.side_effect = lambda n: b"x" * n if n < 1000 else b"x" * 100
# 改成返回固定 100 字节,直到第 N 次后返回空
call_count = [0]
assert size == len(test_content)
def mock_read(size):
call_count[0] += 1
if call_count[0] > 10:
return b""
return b"x" * 100
def test_download_mime_type_rejected(self, tmp_path):
test_content = b"GIF89a" + b"\x00" * 50
dest = str(tmp_path / "test.gif")
mock_resp.read = mock_read
mock_opener.return_value.open.return_value = mock_resp
with self.assertRaises(UrlSecurityError):
mock_resp = MagicMock()
mock_resp.headers = {"Content-Type": "image/gif"}
mock_resp.read.side_effect = [test_content, b""]
with patch("urllib.request.build_opener") as mock_build:
mock_opener = MagicMock()
mock_opener.open.return_value = mock_resp
mock_build.return_value = mock_opener
with pytest.raises(UrlSecurityError, match="Content-Type"):
safe_download_file(
"https://example.com/test",
"https://example.com/test.gif",
dest,
purpose="test",
max_size=500, # 500 字节上限
allowed_mime_types={"image/png"},
)
def test_safe_download_bytes_returns_content(self):
"""safe_download_bytes 应该返回文件内容."""
test_data = b"ID3hello world test audio"
with unittest.mock.patch("urllib.request.build_opener") as mock_opener:
mock_resp = unittest.mock.MagicMock()
mock_resp.headers = {"Content-Type": "audio/mpeg"}
call_count = [0]
def test_download_size_limit_exceeded(self, tmp_path):
"""流式下载时超过大小限制被中断(无 Content-Length header"""
dest = str(tmp_path / "big.bin")
chunk = b"x" * 1024 # 1KB chunks
def mock_read(size):
call_count[0] += 1
if call_count[0] > 1:
return b""
return test_data
mock_resp = MagicMock()
# 没有 Content-Length header,走流式检查
mock_resp.headers = {"Content-Type": "application/octet-stream"}
# 模拟多次读取,超过 5KB 限制(6个chunk = 6KB
mock_resp.read.side_effect = [chunk] * 6 + [b""]
with patch("urllib.request.build_opener") as mock_build:
mock_opener = MagicMock()
mock_opener.open.return_value = mock_resp
mock_build.return_value = mock_opener
with pytest.raises(UrlSecurityError, match="超过大小限制"):
safe_download_file(
"https://example.com/big.bin",
dest,
purpose="test",
max_size=5000, # 5KB limit
)
def test_download_content_length_too_large(self, tmp_path):
dest = str(tmp_path / "big.bin")
mock_resp = MagicMock()
mock_resp.headers = {"Content-Type": "application/octet-stream", "Content-Length": "1000000"}
mock_resp.read.side_effect = [b"data"]
with patch("urllib.request.build_opener") as mock_build:
mock_opener = MagicMock()
mock_opener.open.return_value = mock_resp
mock_build.return_value = mock_opener
with pytest.raises(UrlSecurityError, match="文件过大"):
safe_download_file(
"https://example.com/big.bin",
dest,
purpose="test",
max_size=500000,
)
def test_download_localhost_rejected(self, tmp_path):
"""内网 URL 在下载前就被拒"""
dest = str(tmp_path / "out.bin")
with pytest.raises(UrlSecurityError):
safe_download_file("http://localhost/file", dest)
def test_download_invalid_scheme_rejected(self, tmp_path):
dest = str(tmp_path / "out.bin")
with pytest.raises(UrlSecurityError, match="scheme"):
safe_download_file("ftp://example.com/file", dest)
# ── safe_download_bytes ───────────────────────────────────────────────────────
class TestSafeDownloadBytes:
"""safe_download_bytes 测试"""
def test_download_returns_bytes(self, tmp_path):
test_content = b"hello bytes download test"
mock_resp = MagicMock()
mock_resp.headers = {"Content-Type": "application/octet-stream"}
mock_resp.read.side_effect = [test_content, b""]
with patch("urllib.request.build_opener") as mock_build:
mock_opener = MagicMock()
mock_opener.open.return_value = mock_resp
mock_build.return_value = mock_opener
mock_resp.read = mock_read
mock_opener.return_value.open.return_value = mock_resp
result = safe_download_bytes(
"https://example.com/test.mp3",
"https://example.com/test.bin",
purpose="test",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
)
self.assertEqual(result, test_data)
assert result == test_content
assert isinstance(result, bytes)
def test_download_unsafe_url_raises(self):
with pytest.raises(UrlSecurityError):
safe_download_bytes("http://127.0.0.1/secret")
# ── 常量导出验证 ─────────────────────────────────────────────────────────────
class TestConstants:
"""模块常量验证"""
def test_allowed_schemes(self):
assert "http" in ALLOWED_SCHEMES
assert "https" in ALLOWED_SCHEMES
assert len(ALLOWED_SCHEMES) == 2
def test_allowed_ports(self):
assert 80 in ALLOWED_PORTS
assert 443 in ALLOWED_PORTS
assert len(ALLOWED_PORTS) == 2
def test_max_url_length(self):
assert MAX_URL_LENGTH == 2048