"""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)