158 lines
4.7 KiB
Python
158 lines
4.7 KiB
Python
import hashlib
|
|
import os
|
|
import shutil
|
|
import tempfile
|
|
import unittest
|
|
from contextlib import nullcontext
|
|
from io import StringIO
|
|
|
|
import requests
|
|
from huggingface_hub import snapshot_download
|
|
|
|
from sglang.srt.utils import kill_process_tree
|
|
from sglang.srt.utils.model_file_verifier import compute_sha256, verify
|
|
from sglang.test.ci.ci_register import register_cuda_ci
|
|
from sglang.test.test_utils import (
|
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
|
DEFAULT_URL_FOR_TEST,
|
|
popen_launch_server,
|
|
)
|
|
|
|
register_cuda_ci(est_time=120, suite="nightly-1-gpu", nightly=True)
|
|
|
|
MODEL_NAME = "Qwen/Qwen3-0.6B"
|
|
|
|
|
|
# ======== Base Test Classes ========
|
|
|
|
|
|
class _FakeModelTestCase(unittest.TestCase):
|
|
|
|
FAKE_FILES = {
|
|
"model.safetensors": b"fake safetensors content " * 100,
|
|
"config.json": b'{"model_type": "llama"}',
|
|
"tokenizer.json": b'{"version": "1.0"}',
|
|
}
|
|
|
|
def setUp(self):
|
|
self.test_dir = tempfile.mkdtemp()
|
|
for filename, content in self.FAKE_FILES.items():
|
|
_create_test_file(self.test_dir, filename, content)
|
|
|
|
def tearDown(self):
|
|
shutil.rmtree(self.test_dir, ignore_errors=True)
|
|
|
|
|
|
class _RealModelTestCase(unittest.TestCase):
|
|
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
cls.original_model_path = snapshot_download(MODEL_NAME)
|
|
|
|
def setUp(self):
|
|
self.test_dir = tempfile.mkdtemp()
|
|
shutil.copytree(self.original_model_path, self.test_dir, dirs_exist_ok=True)
|
|
|
|
def tearDown(self):
|
|
shutil.rmtree(self.test_dir, ignore_errors=True)
|
|
|
|
|
|
# ======== Unit Tests ========
|
|
|
|
|
|
class TestModelFileVerifier(_FakeModelTestCase):
|
|
|
|
def test_compute_sha256(self):
|
|
test_file = os.path.join(self.test_dir, "test.bin")
|
|
content = b"hello world"
|
|
with open(test_file, "wb") as f:
|
|
f.write(content)
|
|
|
|
result = compute_sha256(file_path=test_file)
|
|
expected = hashlib.sha256(content).hexdigest()
|
|
self.assertEqual(result, expected)
|
|
|
|
|
|
# ======== HuggingFace Tests ========
|
|
|
|
|
|
class TestModelFileVerifierHF(_RealModelTestCase):
|
|
|
|
def test_verify_with_hf_checksums_source(self):
|
|
verify(model_path=self.test_dir, checksums_source=MODEL_NAME)
|
|
|
|
|
|
# ======== Real Model E2E Tests ========
|
|
|
|
|
|
class TestModelFileVerifierWithRealModel(_RealModelTestCase):
|
|
|
|
def _run_server_test(self, *, corrupt_weights: bool):
|
|
corrupted_file = None
|
|
if corrupt_weights:
|
|
safetensors_files = [
|
|
f for f in os.listdir(self.test_dir) if f.endswith(".safetensors")
|
|
]
|
|
self.assertTrue(len(safetensors_files) > 0, "No safetensors files found")
|
|
corrupted_file = safetensors_files[0]
|
|
_flip_bit_in_file(os.path.join(self.test_dir, corrupted_file))
|
|
|
|
stdout_io, stderr_io = StringIO(), StringIO()
|
|
ctx = self.assertRaises(Exception) if corrupt_weights else nullcontext()
|
|
with ctx:
|
|
process = popen_launch_server(
|
|
model=self.test_dir,
|
|
base_url=DEFAULT_URL_FOR_TEST,
|
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
|
other_args=["--model-checksum", MODEL_NAME],
|
|
return_stdout_stderr=(stdout_io, stderr_io),
|
|
)
|
|
|
|
if corrupt_weights:
|
|
output = stdout_io.getvalue() + stderr_io.getvalue()
|
|
self.assertIn(corrupted_file, output)
|
|
self.assertIn("mismatch", output.lower())
|
|
else:
|
|
try:
|
|
response = requests.post(
|
|
f"{DEFAULT_URL_FOR_TEST}/generate",
|
|
json={"text": "Hello", "sampling_params": {"max_new_tokens": 8}},
|
|
)
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertIn("text", response.json())
|
|
finally:
|
|
kill_process_tree(process.pid)
|
|
|
|
def test_server_launch_with_checksum_intact(self):
|
|
self._run_server_test(corrupt_weights=False)
|
|
|
|
def test_server_launch_fails_with_corrupted_weights(self):
|
|
self._run_server_test(corrupt_weights=True)
|
|
|
|
|
|
# ======== Test Utilities ========
|
|
|
|
|
|
def _create_test_file(directory: str, filename: str, content: bytes) -> str:
|
|
path = os.path.join(directory, filename)
|
|
with open(path, "wb") as f:
|
|
f.write(content)
|
|
return path
|
|
|
|
|
|
def _flip_bit_in_file(file_path: str, byte_offset: int = 100, bit_position: int = 0):
|
|
file_size = os.path.getsize(file_path)
|
|
assert (
|
|
byte_offset < file_size
|
|
), f"byte_offset {byte_offset} >= file_size {file_size}"
|
|
|
|
with open(file_path, "r+b") as f:
|
|
f.seek(byte_offset)
|
|
original_byte = f.read(1)[0]
|
|
f.seek(byte_offset)
|
|
f.write(bytes([original_byte ^ (1 << bit_position)]))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|