c6aac862b1
CI/CD Pipeline / Unit Tests (push) Successful in 1m52s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m17s
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 2m22s
CI/CD Pipeline / Build Production API Image (push) Has been skipped
CI/CD Pipeline / Build Production Worker Image (push) Has been skipped
CI/CD Pipeline / Build Production Web Image (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 1m10s
CI/CD Pipeline / Build Staging Web Image (push) Successful in 1m50s
CI/CD Pipeline / Build Staging API Image (push) Successful in 5m55s
CI/CD Pipeline / Build Staging Worker Image (push) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
- generation.py: 删除重复的shared.url_security import(保留video_processing向后兼容层) - test_url_security.py: mock数据添加ID3魔数头,适配新增的文件魔数校验
297 lines
12 KiB
Python
Executable File
297 lines
12 KiB
Python
Executable File
"""URL 安全校验工具单元测试 — SSRF 防护."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import os
|
||
import sys
|
||
import unittest
|
||
|
||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker"))
|
||
|
||
import shutil
|
||
import tempfile
|
||
|
||
from video_processing.url_security import ( # noqa: E402
|
||
ALLOWED_AUDIO_MIME_TYPES,
|
||
UrlSecurityError,
|
||
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"))
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main()
|
||
|
||
|
||
class TestSafeDownload(unittest.TestCase):
|
||
"""安全下载函数测试."""
|
||
|
||
def setUp(self):
|
||
self.temp_dir = tempfile.mkdtemp()
|
||
|
||
def tearDown(self):
|
||
shutil.rmtree(self.temp_dir, ignore_errors=True)
|
||
|
||
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_safe_download_bytes_rejects_ssrf(self):
|
||
"""SSRF 风险 URL 应该被拒绝下载(bytes 版本)."""
|
||
with self.assertRaises(UrlSecurityError):
|
||
safe_download_bytes("http://localhost/test", purpose="test")
|
||
|
||
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=500,content-length=1000 应被拒绝
|
||
with self.assertRaises(UrlSecurityError):
|
||
safe_download_file(
|
||
"https://example.com/test",
|
||
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,
|
||
)
|
||
|
||
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",
|
||
dest,
|
||
purpose="test",
|
||
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
|
||
)
|
||
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]
|
||
|
||
def mock_read(size):
|
||
call_count[0] += 1
|
||
if call_count[0] > 10:
|
||
return b""
|
||
return b"x" * 100
|
||
|
||
mock_resp.read = mock_read
|
||
mock_opener.return_value.open.return_value = mock_resp
|
||
with self.assertRaises(UrlSecurityError):
|
||
safe_download_file(
|
||
"https://example.com/test",
|
||
dest,
|
||
purpose="test",
|
||
max_size=500, # 500 字节上限
|
||
)
|
||
|
||
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 mock_read(size):
|
||
call_count[0] += 1
|
||
if call_count[0] > 1:
|
||
return b""
|
||
return test_data
|
||
|
||
mock_resp.read = mock_read
|
||
mock_opener.return_value.open.return_value = mock_resp
|
||
result = safe_download_bytes(
|
||
"https://example.com/test.mp3",
|
||
purpose="test",
|
||
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
|
||
)
|
||
self.assertEqual(result, test_data)
|