feat: add safeguards for remote media URLs (#34892)

This commit is contained in:
Mick
2026-08-15 18:12:15 +08:00
committed by GitHub
parent 0c072235f4
commit 35cefd1c51
14 changed files with 453 additions and 45 deletions
@@ -87,6 +87,8 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
server_args.mm_preprocess_cache_size_mb = None
server_args.tokenizer_worker_num = 1
server_args.trust_mm_content_hashes = False
server_args.allowed_media_domains = []
server_args.media_url_max_file_size_mb = 64
hf_config = MagicMock()
mock_hf_processor = MagicMock()
@@ -178,6 +180,8 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
tokenizer_worker_num=1,
base_gpu_id=2,
tp_size=8,
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
@staticmethod
@@ -773,6 +777,8 @@ class TestDoubleBosGuard(CustomTestCase):
server_args.mm_preprocess_cache_size_mb = None
server_args.tokenizer_worker_num = 1
server_args.trust_mm_content_hashes = False
server_args.allowed_media_domains = []
server_args.media_url_max_file_size_mb = 64
mock_hf_processor = MagicMock()
mock_hf_processor.__class__.__name__ = "TestProcessor"
@@ -599,6 +599,8 @@ def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls
mm_processor_worker_num=0,
tokenizer_worker_num=1,
base_gpu_id=0,
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
processor = processor_cls(
hf_config=SimpleNamespace(media_placeholder_token_id=42),
@@ -80,6 +80,8 @@ def make_processor(config, image_processor_cls=None):
mm_processor_worker_num=1,
tokenizer_worker_num=1,
base_gpu_id=0,
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
return QwenVLImageProcessor(
hf_config, server_args, processor, None, skip_mm_pool=True
@@ -0,0 +1,183 @@
"""Security tests for client-supplied remote multimodal media URLs."""
import http.server
import threading
import unittest
from unittest.mock import patch
import requests
from sglang.srt.utils.common import (
_normalize_video_input,
configure_media_url_security,
download_remote_media,
get_image_bytes,
load_audio,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
class _MediaHandler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
if self.path == "/media":
payload = b"remote-media"
self.send_response(200)
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)
return
if self.path == "/same-host-redirect":
self.send_response(302)
self.send_header("Location", "/media")
self.end_headers()
return
if self.path == "/other-host-redirect":
self.send_response(302)
self.send_header(
"Location",
f"http://localhost:{self.server.server_port}/redirect-target",
)
self.end_headers()
return
if self.path == "/redirect-target":
self.server.redirect_target_reached = True
self.send_response(200)
self.end_headers()
self.wfile.write(b"must-not-be-fetched")
return
if self.path == "/oversized":
self.send_response(200)
self.send_header("Content-Length", str(2 * 1024 * 1024))
self.end_headers()
return
if self.path == "/chunked-oversized":
self.send_response(200)
self.end_headers()
self.wfile.write(b"x" * (1024 * 1024 + 1))
return
if self.path == "/redirect-loop":
self.send_response(302)
self.send_header("Location", "/redirect-loop")
self.end_headers()
return
self.send_response(404)
self.end_headers()
def log_message(self, *_):
pass
class TestMediaURLSecurity(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.server = http.server.ThreadingHTTPServer(("127.0.0.1", 0), _MediaHandler)
cls.server.redirect_target_reached = False
cls.thread = threading.Thread(target=cls.server.serve_forever, daemon=True)
cls.thread.start()
cls.port = cls.server.server_port
@classmethod
def tearDownClass(cls):
cls.server.shutdown()
cls.server.server_close()
cls.thread.join()
def setUp(self):
self.server.redirect_target_reached = False
configure_media_url_security([], max_file_size_mb=64)
def tearDown(self):
configure_media_url_security([], max_file_size_mb=64)
def _url(self, path, host="127.0.0.1"):
return f"http://{host}:{self.port}{path}"
def test_unrestricted_mode_preserves_remote_media_compatibility(self):
self.assertEqual(
download_remote_media(self._url("/media", host="localhost"), timeout=5),
b"remote-media",
)
def test_exact_domain_allowlist(self):
configure_media_url_security(["127.0.0.1"], max_file_size_mb=64)
self.assertEqual(
download_remote_media(self._url("/media"), timeout=5), b"remote-media"
)
with self.assertRaisesRegex(ValueError, "not allowed"):
download_remote_media(self._url("/media", host="localhost"), timeout=5)
with self.assertRaisesRegex(ValueError, "not allowed"):
download_remote_media("http://169.254.169.254/latest/meta-data", timeout=5)
def test_redirect_destination_is_checked_before_fetch(self):
configure_media_url_security(["127.0.0.1"], max_file_size_mb=64)
with self.assertRaisesRegex(ValueError, "not allowed"):
download_remote_media(self._url("/other-host-redirect"), timeout=5)
self.assertFalse(self.server.redirect_target_reached)
def test_same_domain_redirect_is_allowed(self):
configure_media_url_security(["127.0.0.1"], max_file_size_mb=64)
self.assertEqual(
download_remote_media(self._url("/same-host-redirect"), timeout=5),
b"remote-media",
)
def test_redirect_count_is_bounded(self):
configure_media_url_security(["127.0.0.1"], max_file_size_mb=64)
with self.assertRaises(requests.exceptions.TooManyRedirects):
download_remote_media(self._url("/redirect-loop"), timeout=5)
def test_declared_oversized_response_is_rejected(self):
configure_media_url_security(["127.0.0.1"], max_file_size_mb=1)
with self.assertRaisesRegex(ValueError, "download limit"):
download_remote_media(self._url("/oversized"), timeout=5)
def test_streamed_oversized_response_is_rejected(self):
configure_media_url_security(["127.0.0.1"], max_file_size_mb=1)
with self.assertRaisesRegex(ValueError, "download limit"):
download_remote_media(self._url("/chunked-oversized"), timeout=5)
def test_invalid_allowlist_entries_are_rejected(self):
for domain in (
"https://media.example.com",
"media.example.com/path",
"media.example.com:443",
"",
):
with self.subTest(domain=domain):
with self.assertRaises(ValueError):
configure_media_url_security([domain], max_file_size_mb=64)
def test_backslash_userinfo_parser_confusion_cannot_bypass_allowlist(self):
configure_media_url_security(["safe.example.org"], max_file_size_mb=64)
with self.assertRaisesRegex(ValueError, "not allowed"):
download_remote_media(
r"https://evil.example\@safe.example.org/media", timeout=5
)
def test_all_common_loaders_share_the_policy(self):
blocked = ValueError("media URL domain is not allowed")
with patch(
"sglang.srt.utils.common.download_remote_media", side_effect=blocked
) as download:
for loader in (
get_image_bytes,
_normalize_video_input,
load_audio,
):
with self.subTest(loader=loader.__name__):
with self.assertRaisesRegex(ValueError, "not allowed"):
loader("https://blocked.example/media")
self.assertEqual(download.call_count, 3)
if __name__ == "__main__":
unittest.main()
@@ -8,6 +8,7 @@ import argparse
import unittest
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import configure_media_url_security
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -79,6 +80,33 @@ class TestServerArgsAnnotatedCli(CustomTestCase):
self.assertEqual(sa.extra_metric_labels, {"k": "v"})
self.assertEqual(sa.forward_hooks, [{"type": "test"}])
def test_media_url_security_args(self):
try:
sa = self._parse(
[
"--allowed-media-domains",
"Media.Example.com.",
"127.0.0.1",
"--media-url-max-file-size-mb",
"32",
]
)
self.assertEqual(
sa.allowed_media_domains, ["127.0.0.1", "media.example.com"]
)
self.assertEqual(sa.media_url_max_file_size_mb, 32)
finally:
configure_media_url_security([], max_file_size_mb=64)
def test_media_url_security_args_reject_invalid_values(self):
try:
with self.assertRaises(ValueError):
self._parse(["--allowed-media-domains", "https://media.example.com"])
with self.assertRaises(ValueError):
self._parse(["--media-url-max-file-size-mb", "-1"])
finally:
configure_media_url_security([], max_file_size_mb=64)
def test_literal_auto_derives_choices(self):
"""Literal type annotations produce argparse choices automatically."""
sa = self._parse(