feat: add safeguards for remote media URLs (#34892)
This commit is contained in:
@@ -56,6 +56,15 @@ You can find all arguments by `python3 -m sglang.launch_server --help`
|
||||
- To enable fp8 weight quantization, add `--quantization fp8` on a fp16 checkpoint or directly load a fp8 checkpoint without specifying any arguments.
|
||||
- To enable fp8 kv cache quantization, add `--kv-cache-dtype fp8_e4m3` or `--kv-cache-dtype fp8_e5m2`.
|
||||
- To enable deterministic inference and batch invariant operations, add `--enable-deterministic-inference`. More details can be found in [deterministic inference document](./deterministic_inference).
|
||||
- If a multimodal server accepts requests from untrusted clients, restrict remote image, video, and audio URLs with `--allowed-media-domains`. SGLang checks the initial URL and every redirect destination against the exact-hostname allowlist. Remote media downloads are limited to 64 MiB by default; adjust `--media-url-max-file-size-mb` when larger trusted media is required.
|
||||
|
||||
```bash Command
|
||||
python -m sglang.launch_server \
|
||||
--model-path Qwen/Qwen2.5-VL-7B-Instruct \
|
||||
--allowed-media-domains upload.wikimedia.org raw.githubusercontent.com
|
||||
```
|
||||
|
||||
Without `--allowed-media-domains`, HTTP(S) media from any domain remains allowed for backward compatibility. Do not expose that configuration to untrusted users. Local paths and `data:` URLs are not governed by the domain allowlist.
|
||||
- To enable decode context parallelism for MLA models, add `--dcp-size N`. See [Decode Context Parallelism](./dcp).
|
||||
- If the model does not have a chat template in the Hugging Face tokenizer, you can specify a [custom chat template](../references/custom_chat_template). If the tokenizer has multiple named templates (e.g., 'default', 'tool_use'), you can select one using `--hf-chat-template-name tool_use`.
|
||||
- To run tensor parallelism on multiple nodes, add `--nnodes 2`. If you have two nodes with two GPUs on each node and want to run TP=4, let `sgl-dev-0` be the hostname of the first node and `50000` be an available port, you can use the following commands. If you meet deadlock, please try to add `--disable-cuda-graph`
|
||||
@@ -3293,6 +3302,18 @@ Please consult the documentation below and [server_args.py](https://github.com/s
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>{}</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Type: JSON / Dict</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--allowed-media-domains`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Restrict client-supplied HTTP(S) media URLs and redirect destinations to these exact hostnames.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Unrestricted</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Space-separated hostnames</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--media-url-max-file-size-mb`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Maximum streamed size in MiB for one remote media download. Set to 0 to disable the limit.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`64`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Type: int</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--mm-enable-dp-encoder`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Enabling data parallelism for mm encoder. The dp size will be set to the tp size automatically.</td>
|
||||
|
||||
@@ -85,6 +85,7 @@ from sglang.srt.utils import (
|
||||
CLIENT_MEDIA_EXCEPTIONS,
|
||||
add_prometheus_middleware,
|
||||
configure_logger,
|
||||
configure_media_url_security,
|
||||
load_audio,
|
||||
load_image,
|
||||
load_video,
|
||||
@@ -308,6 +309,10 @@ class MMEncoder:
|
||||
argument."""
|
||||
logger.info(f"init MMEncoder {rank}/{server_args.tp_size}")
|
||||
self.server_args = server_args
|
||||
configure_media_url_security(
|
||||
server_args.allowed_media_domains,
|
||||
server_args.media_url_max_file_size_mb,
|
||||
)
|
||||
publish(server_args, role="encoder")
|
||||
self.rank = rank
|
||||
# DP rank for metric labels; overridden by run_dp_worker in DP mode.
|
||||
|
||||
@@ -32,6 +32,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
CLIENT_MEDIA_EXCEPTIONS,
|
||||
configure_media_url_security,
|
||||
envs,
|
||||
is_cpu,
|
||||
is_npu,
|
||||
@@ -199,6 +200,10 @@ class BaseMultimodalProcessor(ABC):
|
||||
self._processor = _processor
|
||||
self.server_args = server_args
|
||||
self.transport_mode = transport_mode
|
||||
configure_media_url_security(
|
||||
server_args.allowed_media_domains,
|
||||
server_args.media_url_max_file_size_mb,
|
||||
)
|
||||
configured_mm_feature_transport = getattr(
|
||||
server_args, "mm_feature_transport", "cpu"
|
||||
)
|
||||
|
||||
@@ -24,7 +24,6 @@ from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import logging
|
||||
import urllib.request
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Dict, List, Optional, Tuple, Type, Union
|
||||
|
||||
@@ -53,6 +52,7 @@ from sglang.srt.parser.inkling_tokenizer import IMAGE_TOKEN_ID as INKLING_IMAGE_
|
||||
from sglang.srt.parser.inkling_tokenizer import (
|
||||
INKLING_SPECIAL_TOKEN_IDS,
|
||||
)
|
||||
from sglang.srt.utils.common import download_remote_media
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -82,8 +82,7 @@ def _resolve_media_item(item):
|
||||
header, _, payload = url.partition(",")
|
||||
return base64.b64decode(payload) if ";base64" in header else payload.encode()
|
||||
if url.startswith(("http://", "https://")):
|
||||
with urllib.request.urlopen(url, timeout=30) as resp:
|
||||
return resp.read()
|
||||
return download_remote_media(url, timeout=30)
|
||||
return url # plain path / file:// -> handled by the per-modality byte loader
|
||||
|
||||
|
||||
|
||||
@@ -186,18 +186,14 @@ class MiMoAudioPipeline:
|
||||
dl_start = time.perf_counter()
|
||||
timeout = int(os.getenv("REQUEST_TIMEOUT", "5"))
|
||||
try:
|
||||
with common.get_mm_http_session().get(
|
||||
audio, stream=True, timeout=timeout
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
|
||||
if dl_elapsed_ms > 1000.0:
|
||||
content_len = len(response.content)
|
||||
logger.warning(
|
||||
f"Slow audio download: {dl_elapsed_ms:.2f}ms, "
|
||||
f"size={content_len / 1024:.1f}KB, url={audio}"
|
||||
)
|
||||
file = io.BytesIO(response.content)
|
||||
content = common.download_remote_media(audio, timeout=timeout)
|
||||
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
|
||||
if dl_elapsed_ms > 1000.0:
|
||||
logger.warning(
|
||||
f"Slow audio download: {dl_elapsed_ms:.2f}ms, "
|
||||
f"size={len(content) / 1024:.1f}KB, url={audio}"
|
||||
)
|
||||
file = io.BytesIO(content)
|
||||
except Exception as e:
|
||||
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
|
||||
logger.error(
|
||||
|
||||
@@ -5,6 +5,7 @@ import base64
|
||||
import copy
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
@@ -13,7 +14,6 @@ from io import BytesIO
|
||||
from typing import List, Literal, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import requests
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from fastapi import HTTPException
|
||||
@@ -40,6 +40,7 @@ from sglang.srt.multimodal.processors.mimo_audio import (
|
||||
)
|
||||
from sglang.srt.multimodal.processors.qwen_vl import smart_nframes
|
||||
from sglang.srt.utils import ImageData, VideoData
|
||||
from sglang.srt.utils.common import download_remote_media
|
||||
from sglang.utils import logger
|
||||
|
||||
|
||||
@@ -485,12 +486,14 @@ class MiMoProcessor:
|
||||
|
||||
@staticmethod
|
||||
def has_audio_track(path_or_data) -> bool:
|
||||
# In-process probe via torchcodec for bytes/path; ffprobe range
|
||||
# request for HTTP URLs so we do not pre-download the blob here.
|
||||
# Never hand a client-supplied URL to ffprobe: its internal HTTP client
|
||||
# would bypass the shared domain and redirect policy. Resolve it through
|
||||
# the guarded downloader first, then probe the resulting bytes in-process.
|
||||
if isinstance(path_or_data, str) and path_or_data.startswith(
|
||||
("http://", "https://")
|
||||
):
|
||||
return _ffprobe_has_audio(path_or_data, stdin=None, label=path_or_data)
|
||||
timeout = int(os.getenv("REQUEST_TIMEOUT", "10"))
|
||||
path_or_data = download_remote_media(path_or_data, timeout=timeout)
|
||||
|
||||
if isinstance(path_or_data, bytes):
|
||||
source = BytesIO(path_or_data)
|
||||
@@ -1446,10 +1449,8 @@ class MiMoProcessor:
|
||||
image_obj = image
|
||||
elif isinstance(image, str):
|
||||
if image.startswith("http://") or image.startswith("https://"):
|
||||
with requests.get(image, stream=True) as response:
|
||||
response.raise_for_status()
|
||||
with BytesIO(response.content) as bio:
|
||||
image_obj = copy.deepcopy(Image.open(bio))
|
||||
with BytesIO(download_remote_media(image, timeout=3)) as bio:
|
||||
image_obj = copy.deepcopy(Image.open(bio))
|
||||
elif image.startswith("file://"):
|
||||
image_obj = Image.open(image[7:])
|
||||
elif image.startswith("data:image"):
|
||||
|
||||
@@ -6,7 +6,6 @@ from typing import Dict, List, Optional, Tuple, Union
|
||||
from urllib.parse import unquote, urlparse
|
||||
|
||||
import pybase64
|
||||
import requests
|
||||
import torch
|
||||
|
||||
from sglang.srt.managers.schedule_batch import (
|
||||
@@ -21,6 +20,7 @@ from sglang.srt.multimodal.processors.base_processor import (
|
||||
from sglang.srt.multimodal.processors.base_processor import (
|
||||
MultimodalSpecialTokens,
|
||||
)
|
||||
from sglang.srt.utils.common import download_remote_media
|
||||
|
||||
|
||||
class MossVLImageProcessor(SGLangBaseProcessor):
|
||||
@@ -426,13 +426,10 @@ class MossVLImageProcessor(SGLangBaseProcessor):
|
||||
|
||||
if value.startswith(("http://", "https://")):
|
||||
timeout = int(os.getenv("REQUEST_TIMEOUT", "10"))
|
||||
response = requests.get(value, stream=True, timeout=timeout)
|
||||
response.raise_for_status()
|
||||
content = download_remote_media(value, timeout=timeout)
|
||||
suffix = os.path.splitext(urlparse(value).path)[1] or ".mp4"
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as f:
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
f.write(content)
|
||||
return f.name, f.name
|
||||
|
||||
if value.startswith("data:"):
|
||||
|
||||
@@ -70,6 +70,7 @@ from sglang.srt.speculative.decoupled_spec_io import DecoupledSpecIpcConfig
|
||||
from sglang.srt.utils.common import (
|
||||
LORA_TARGET_ALL_MODULES,
|
||||
SUPPORTED_LORA_TARGET_MODULES,
|
||||
configure_media_url_security,
|
||||
get_device,
|
||||
get_device_memory_capacity,
|
||||
get_device_sm,
|
||||
@@ -2764,6 +2765,19 @@ class ServerArgs:
|
||||
"environment override when this argument is 0.",
|
||||
NS("mm"),
|
||||
] = 0
|
||||
allowed_media_domains: A[
|
||||
List[str],
|
||||
"Restrict client-supplied HTTP(S) image, video, and audio URLs to these "
|
||||
"exact hostnames. Redirect destinations are checked against the same "
|
||||
"allowlist. When unset, remote media from any domain is allowed.",
|
||||
NS("mm"),
|
||||
] = dataclasses.field(default_factory=list)
|
||||
media_url_max_file_size_mb: A[
|
||||
int,
|
||||
"Maximum size in MiB for one client-supplied remote media download. "
|
||||
"The limit is enforced while streaming; set to 0 to disable it.",
|
||||
NS("mm"),
|
||||
] = 64
|
||||
mm_preprocess_cache_size_mb: A[
|
||||
Optional[int],
|
||||
"CPU memory budget for content-addressed multimodal preprocessing "
|
||||
@@ -3561,6 +3575,7 @@ class ServerArgs:
|
||||
|
||||
self._handle_moe_runner_backend_alias()
|
||||
self._handle_return_hidden_states_mode()
|
||||
self._handle_media_url_security()
|
||||
if self.model_path.lower() in ["none", "dummy"]:
|
||||
return
|
||||
|
||||
@@ -4051,6 +4066,13 @@ class ServerArgs:
|
||||
f"but got {type(self.mm_process_config[key])}"
|
||||
)
|
||||
|
||||
def _handle_media_url_security(self):
|
||||
"""Normalize and publish the media URL policy before workers start."""
|
||||
self.allowed_media_domains = configure_media_url_security(
|
||||
self.allowed_media_domains,
|
||||
self.media_url_max_file_size_mb,
|
||||
)
|
||||
|
||||
def _handle_deprecated_args(self):
|
||||
if self.disable_fast_image_processor:
|
||||
if self.image_processor_backend not in {"auto", "pil"}:
|
||||
|
||||
@@ -25,6 +25,7 @@ import gc
|
||||
import importlib
|
||||
import inspect
|
||||
import io
|
||||
import ipaddress
|
||||
import itertools
|
||||
import json
|
||||
import logging
|
||||
@@ -75,7 +76,7 @@ from typing import (
|
||||
)
|
||||
from unittest import SkipTest
|
||||
from unittest.case import _ShouldStop
|
||||
from urllib.parse import unquote, urlparse
|
||||
from urllib.parse import unquote, urljoin, urlparse
|
||||
|
||||
import numpy as np
|
||||
import orjson
|
||||
@@ -1510,6 +1511,158 @@ def set_random_seed(seed: int) -> None:
|
||||
|
||||
_mm_http_session = threading.local()
|
||||
|
||||
_DEFAULT_MEDIA_URL_MAX_FILE_SIZE_MB = 64
|
||||
_MAX_MEDIA_URL_REDIRECTS = 5
|
||||
_MEDIA_URL_REDIRECT_STATUS_CODES = {301, 302, 303, 307, 308}
|
||||
_allowed_media_domains: frozenset[str] = frozenset()
|
||||
_media_url_max_file_size_bytes = _DEFAULT_MEDIA_URL_MAX_FILE_SIZE_MB * 1024 * 1024
|
||||
|
||||
|
||||
def _normalize_media_domain(domain: str) -> str:
|
||||
if not isinstance(domain, str):
|
||||
raise ValueError("allowed media domains must be strings")
|
||||
|
||||
domain = domain.strip().rstrip(".")
|
||||
if not domain:
|
||||
raise ValueError("allowed media domains cannot be empty")
|
||||
if "://" in domain or any(char in domain for char in "/?#@"):
|
||||
raise ValueError(
|
||||
f"Invalid allowed media domain {domain!r}: provide a hostname only"
|
||||
)
|
||||
|
||||
# Brackets are URL syntax, not part of an IPv6 hostname.
|
||||
if domain.startswith("[") and domain.endswith("]"):
|
||||
domain = domain[1:-1]
|
||||
try:
|
||||
return str(ipaddress.ip_address(domain))
|
||||
except ValueError:
|
||||
if ":" in domain:
|
||||
raise ValueError(
|
||||
f"Invalid allowed media domain {domain!r}: ports are not supported"
|
||||
)
|
||||
|
||||
try:
|
||||
normalized = domain.encode("idna").decode("ascii").lower()
|
||||
except UnicodeError as e:
|
||||
raise ValueError(f"Invalid allowed media domain {domain!r}") from e
|
||||
if not normalized:
|
||||
raise ValueError("allowed media domains cannot be empty")
|
||||
return normalized
|
||||
|
||||
|
||||
def configure_media_url_security(
|
||||
allowed_media_domains: Optional[Sequence[str]] = None,
|
||||
max_file_size_mb: int = _DEFAULT_MEDIA_URL_MAX_FILE_SIZE_MB,
|
||||
) -> list[str]:
|
||||
"""Configure process-wide safeguards for client-supplied media URLs.
|
||||
|
||||
A serving worker hosts one engine configuration, while media loading fans
|
||||
out to worker threads. Keeping the immutable policy here makes the same
|
||||
checks apply to image, video, audio, cache, and model-specific loaders.
|
||||
"""
|
||||
|
||||
if max_file_size_mb < 0:
|
||||
raise ValueError("media_url_max_file_size_mb must be non-negative")
|
||||
|
||||
normalized_domains = sorted(
|
||||
{_normalize_media_domain(domain) for domain in allowed_media_domains or []}
|
||||
)
|
||||
global _allowed_media_domains, _media_url_max_file_size_bytes
|
||||
_allowed_media_domains = frozenset(normalized_domains)
|
||||
_media_url_max_file_size_bytes = max_file_size_mb * 1024 * 1024
|
||||
return normalized_domains
|
||||
|
||||
|
||||
def _assert_media_url_allowed(url: str) -> None:
|
||||
parsed = urlparse(url)
|
||||
if parsed.scheme not in {"http", "https"} or parsed.hostname is None:
|
||||
raise ValueError(f"Invalid media URL: {url!r}")
|
||||
|
||||
hostname = _normalize_media_domain(parsed.hostname)
|
||||
if _allowed_media_domains and hostname not in _allowed_media_domains:
|
||||
raise ValueError(
|
||||
"Media URL domain is not allowed. "
|
||||
f"Allowed domains: {sorted(_allowed_media_domains)}; "
|
||||
f"input domain: {hostname}"
|
||||
)
|
||||
|
||||
|
||||
def download_remote_media(url: str, timeout: float) -> bytes:
|
||||
"""Download one HTTP(S) media object under the configured URL policy.
|
||||
|
||||
Redirects are followed manually so every destination is validated before
|
||||
a connection is made. The response is streamed to enforce both the total
|
||||
request deadline and the configured byte limit without first buffering an
|
||||
attacker-controlled body in memory.
|
||||
"""
|
||||
|
||||
if timeout <= 0:
|
||||
raise ValueError("media URL timeout must be positive")
|
||||
|
||||
session = get_mm_http_session()
|
||||
deadline = time.monotonic() + timeout
|
||||
current_url = url
|
||||
|
||||
for redirect_count in range(_MAX_MEDIA_URL_REDIRECTS + 1):
|
||||
# Validate the same normalized URL representation that requests sends
|
||||
# to urllib3. This avoids parser disagreements around backslashes and
|
||||
# userinfo separators.
|
||||
prepared_url = requests.Request("GET", current_url).prepare().url
|
||||
if prepared_url is None:
|
||||
raise ValueError(f"Invalid media URL: {current_url!r}")
|
||||
_assert_media_url_allowed(prepared_url)
|
||||
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
raise requests.exceptions.Timeout(
|
||||
f"Timed out while downloading media URL: {url}"
|
||||
)
|
||||
|
||||
with session.get(
|
||||
prepared_url,
|
||||
allow_redirects=False,
|
||||
stream=True,
|
||||
timeout=remaining,
|
||||
) as response:
|
||||
location = response.headers.get("Location")
|
||||
if response.status_code in _MEDIA_URL_REDIRECT_STATUS_CODES and location:
|
||||
if redirect_count == _MAX_MEDIA_URL_REDIRECTS:
|
||||
raise requests.exceptions.TooManyRedirects(
|
||||
f"Media URL exceeded {_MAX_MEDIA_URL_REDIRECTS} redirects: {url}"
|
||||
)
|
||||
current_url = urljoin(response.url, location)
|
||||
continue
|
||||
|
||||
response.raise_for_status()
|
||||
max_bytes = _media_url_max_file_size_bytes
|
||||
content_length = response.headers.get("Content-Length")
|
||||
if max_bytes and content_length is not None:
|
||||
try:
|
||||
declared_size = int(content_length)
|
||||
except ValueError:
|
||||
declared_size = None
|
||||
if declared_size is not None and declared_size > max_bytes:
|
||||
raise ValueError(
|
||||
f"Remote media exceeds the {max_bytes} byte download limit"
|
||||
)
|
||||
|
||||
content = bytearray()
|
||||
for chunk in response.iter_content(chunk_size=64 * 1024):
|
||||
if not chunk:
|
||||
continue
|
||||
if time.monotonic() > deadline:
|
||||
raise requests.exceptions.Timeout(
|
||||
f"Timed out while downloading media URL: {url}"
|
||||
)
|
||||
if max_bytes and len(content) + len(chunk) > max_bytes:
|
||||
raise ValueError(
|
||||
f"Remote media exceeds the {max_bytes} byte download limit"
|
||||
)
|
||||
content.extend(chunk)
|
||||
return bytes(content)
|
||||
|
||||
raise AssertionError("unreachable")
|
||||
|
||||
|
||||
def get_mm_http_session() -> requests.Session:
|
||||
"""Per-thread HTTP session for multimodal downloads, to pool/reuse TCP
|
||||
@@ -1548,9 +1701,7 @@ def load_audio(
|
||||
audio_file.startswith("http://") or audio_file.startswith("https://")
|
||||
):
|
||||
timeout = int(os.getenv("REQUEST_TIMEOUT", "5"))
|
||||
with get_mm_http_session().get(audio_file, timeout=timeout) as response:
|
||||
response.raise_for_status()
|
||||
source = response.content
|
||||
source = download_remote_media(audio_file, timeout=timeout)
|
||||
elif isinstance(audio_file, str) and audio_file.startswith("file://"):
|
||||
source = unquote(urlparse(audio_file).path)
|
||||
elif isinstance(audio_file, str):
|
||||
@@ -1753,13 +1904,7 @@ def get_image_bytes(image_file: Union[str, bytes]) -> bytes:
|
||||
return image_file
|
||||
if image_file.startswith(("http://", "https://")):
|
||||
timeout = int(os.getenv("REQUEST_TIMEOUT", "3"))
|
||||
response = get_mm_http_session().get(image_file, timeout=timeout)
|
||||
try:
|
||||
response.raise_for_status()
|
||||
result = response.content
|
||||
finally:
|
||||
response.close()
|
||||
return result
|
||||
return download_remote_media(image_file, timeout=timeout)
|
||||
if image_file.startswith(("file://", "/")):
|
||||
with open(image_file, "rb") as f:
|
||||
return f.read()
|
||||
@@ -1785,11 +1930,7 @@ def _normalize_video_input(
|
||||
elif isinstance(video_file, str):
|
||||
if video_file.startswith(("http://", "https://")):
|
||||
timeout = int(os.getenv("REQUEST_TIMEOUT", "10"))
|
||||
with get_mm_http_session().get(
|
||||
video_file, stream=True, timeout=timeout
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
return response.content
|
||||
return download_remote_media(video_file, timeout=timeout)
|
||||
elif video_file.startswith("data:"):
|
||||
_, encoded = video_file.split(",", 1)
|
||||
return pybase64.b64decode(encoded, validate=True)
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user