Clean up ServerArgs post-init dispatch (#30186)

This commit is contained in:
Lianmin Zheng
2026-07-05 23:05:07 -07:00
committed by GitHub
parent 5eb1b6a7ba
commit 24c42c90be
3 changed files with 147 additions and 122 deletions
@@ -11,41 +11,43 @@ register_amd_ci(est_time=1, suite="stage-b-test-1-gpu-small-amd")
class TestMmProcessConfigValidation(unittest.TestCase):
"""Server-args validation for mm_process_config."""
def _validate_config(self, mm_process_config):
args = ServerArgs(model_path="dummy", mm_process_config=mm_process_config)
args._handle_multimodal()
return args
def test_valid_config_accepted(self):
args = ServerArgs(
model_path="dummy",
mm_process_config={"image": {"max_pixels": 5000000}},
)
args = self._validate_config({"image": {"max_pixels": 5000000}})
self.assertEqual(args.mm_process_config, {"image": {"max_pixels": 5000000}})
def test_empty_config_accepted(self):
args = ServerArgs(model_path="dummy", mm_process_config={})
args = self._validate_config({})
self.assertEqual(args.mm_process_config, {})
def test_none_config_defaults_to_empty_dict(self):
args = ServerArgs(model_path="dummy", mm_process_config=None)
args = self._validate_config(None)
# None is kept as-is for dummy models (default happens after early return)
# but for real models it would be set to {}
self.assertIsNone(args.mm_process_config)
def test_top_level_non_dict_rejected(self):
with self.assertRaises(TypeError) as ctx:
ServerArgs(model_path="dummy", mm_process_config="bad")
self._validate_config("bad")
self.assertIn("mm_process_config must be a dict", str(ctx.exception))
def test_modality_non_dict_rejected_image(self):
with self.assertRaises(TypeError) as ctx:
ServerArgs(model_path="dummy", mm_process_config={"image": "bad"})
self._validate_config({"image": "bad"})
self.assertIn("mm_process_config['image'] must be a dict", str(ctx.exception))
def test_modality_non_dict_rejected_video(self):
with self.assertRaises(TypeError) as ctx:
ServerArgs(model_path="dummy", mm_process_config={"video": 123})
self._validate_config({"video": 123})
self.assertIn("mm_process_config['video'] must be a dict", str(ctx.exception))
def test_modality_non_dict_rejected_audio(self):
with self.assertRaises(TypeError) as ctx:
ServerArgs(model_path="dummy", mm_process_config={"audio": [1, 2]})
self._validate_config({"audio": [1, 2]})
self.assertIn("mm_process_config['audio'] must be a dict", str(ctx.exception))
def test_multi_modality_config_accepted(self):
@@ -54,7 +56,7 @@ class TestMmProcessConfigValidation(unittest.TestCase):
"video": {"max_pixels": 602112},
"audio": {"sample_rate": 16000},
}
args = ServerArgs(model_path="dummy", mm_process_config=config)
args = self._validate_config(config)
self.assertEqual(args.mm_process_config, config)
@@ -108,27 +108,34 @@ class TestMambaCacheStochasticRounding(unittest.TestCase):
class TestLoadBalanceMethod(unittest.TestCase):
def _load_balance_args(self, **kwargs):
server_args = ServerArgs(model_path="dummy", **kwargs)
server_args._handle_pd_disaggregation()
server_args._handle_load_balance_method()
return server_args
def test_non_pd_defaults_to_round_robin(self):
server_args = ServerArgs(model_path="dummy", disaggregation_mode="null")
server_args = self._load_balance_args(disaggregation_mode="null")
self.assertEqual(server_args.load_balance_method, "round_robin")
def test_pd_prefill_defaults_to_follow_bootstrap_room(self):
server_args = ServerArgs(model_path="dummy", disaggregation_mode="prefill")
server_args = self._load_balance_args(disaggregation_mode="prefill")
self.assertEqual(server_args.load_balance_method, "follow_bootstrap_room")
def test_pd_decode_defaults_to_round_robin(self):
server_args = ServerArgs(model_path="dummy", disaggregation_mode="decode")
server_args = self._load_balance_args(disaggregation_mode="decode")
self.assertEqual(server_args.load_balance_method, "round_robin")
def test_pd_decode_radix_cache_rejects_hisparse(self):
server_args = ServerArgs(
model_path="dummy",
disaggregation_mode="decode",
disaggregation_decode_enable_radix_cache=True,
disaggregation_transfer_backend="nixl",
enable_hisparse=True,
)
with self.assertRaises(ValueError) as context:
ServerArgs(
model_path="dummy",
disaggregation_mode="decode",
disaggregation_decode_enable_radix_cache=True,
disaggregation_transfer_backend="nixl",
enable_hisparse=True,
)
server_args._handle_pd_disaggregation()
self.assertIn(
"--disaggregation-decode-enable-radix-cache is incompatible with "
@@ -137,8 +144,7 @@ class TestLoadBalanceMethod(unittest.TestCase):
)
def test_pd_decode_radix_cache_allows_mooncake(self):
server_args = ServerArgs(
model_path="dummy",
server_args = self._load_balance_args(
disaggregation_mode="decode",
disaggregation_decode_enable_radix_cache=True,
disaggregation_transfer_backend="mooncake",
@@ -147,13 +153,14 @@ class TestLoadBalanceMethod(unittest.TestCase):
self.assertFalse(server_args.disable_radix_cache)
def test_pd_decode_radix_cache_rejects_fake_backend(self):
server_args = ServerArgs(
model_path="dummy",
disaggregation_mode="decode",
disaggregation_decode_enable_radix_cache=True,
disaggregation_transfer_backend="fake",
)
with self.assertRaises(ValueError) as context:
ServerArgs(
model_path="dummy",
disaggregation_mode="decode",
disaggregation_decode_enable_radix_cache=True,
disaggregation_transfer_backend="fake",
)
server_args._handle_pd_disaggregation()
self.assertIn(
"--disaggregation-decode-enable-radix-cache is incompatible "
@@ -162,8 +169,7 @@ class TestLoadBalanceMethod(unittest.TestCase):
)
def test_pd_decode_radix_cache_allows_ascend(self):
server_args = ServerArgs(
model_path="dummy",
server_args = self._load_balance_args(
disaggregation_mode="decode",
disaggregation_decode_enable_radix_cache=True,
disaggregation_transfer_backend="ascend",
@@ -172,8 +178,7 @@ class TestLoadBalanceMethod(unittest.TestCase):
self.assertFalse(server_args.disable_radix_cache)
def test_pd_decode_radix_cache_allows_mooncake_tcp(self):
server_args = ServerArgs(
model_path="dummy",
server_args = self._load_balance_args(
disaggregation_mode="decode",
disaggregation_decode_enable_radix_cache=True,
disaggregation_transfer_backend="mooncake_tcp",
@@ -709,6 +714,11 @@ class TestPortArgs(unittest.TestCase):
class TestSSLArgs(unittest.TestCase):
def _validate_ssl(self, **kwargs):
server_args = ServerArgs(model_path="dummy", **kwargs)
server_args._handle_ssl_validation()
return server_args
def test_default_ssl_fields_are_none(self):
server_args = ServerArgs(model_path="dummy")
self.assertIsNone(server_args.ssl_keyfile)
@@ -718,19 +728,17 @@ class TestSSLArgs(unittest.TestCase):
def test_ssl_keyfile_without_certfile_raises(self):
with self.assertRaises(ValueError) as context:
ServerArgs(model_path="dummy", ssl_keyfile="key.pem")
self._validate_ssl(ssl_keyfile="key.pem")
self.assertIn("--ssl-certfile", str(context.exception))
def test_ssl_certfile_without_keyfile_raises(self):
with self.assertRaises(ValueError) as context:
ServerArgs(model_path="dummy", ssl_certfile="cert.pem")
self._validate_ssl(ssl_certfile="cert.pem")
self.assertIn("--ssl-keyfile", str(context.exception))
@patch("os.path.isfile", return_value=True)
def test_ssl_both_keyfile_and_certfile_accepted(self, _mock_isfile):
server_args = ServerArgs(
model_path="dummy", ssl_keyfile="key.pem", ssl_certfile="cert.pem"
)
server_args = self._validate_ssl(ssl_keyfile="key.pem", ssl_certfile="cert.pem")
self.assertEqual(server_args.ssl_keyfile, "key.pem")
self.assertEqual(server_args.ssl_certfile, "cert.pem")
@@ -748,9 +756,7 @@ class TestSSLArgs(unittest.TestCase):
@patch("os.path.isfile", return_value=True)
def test_url_returns_https_with_ssl(self, _mock_isfile):
server_args = ServerArgs(
model_path="dummy", ssl_keyfile="key.pem", ssl_certfile="cert.pem"
)
server_args = self._validate_ssl(ssl_keyfile="key.pem", ssl_certfile="cert.pem")
self.assertTrue(server_args.url().startswith("https://"))
@patch("os.path.isfile", return_value=True)
@@ -780,15 +786,12 @@ class TestSSLArgs(unittest.TestCase):
@patch("os.path.isfile", return_value=True)
def test_ssl_verify_with_ssl_no_ca(self, _mock_isfile):
server_args = ServerArgs(
model_path="dummy", ssl_keyfile="key.pem", ssl_certfile="cert.pem"
)
server_args = self._validate_ssl(ssl_keyfile="key.pem", ssl_certfile="cert.pem")
self.assertIs(server_args.ssl_verify(), False)
@patch("os.path.isfile", return_value=True)
def test_ssl_verify_with_ssl_and_ca(self, _mock_isfile):
server_args = ServerArgs(
model_path="dummy",
server_args = self._validate_ssl(
ssl_keyfile="key.pem",
ssl_certfile="cert.pem",
ssl_ca_certs="ca.pem",
@@ -797,18 +800,17 @@ class TestSSLArgs(unittest.TestCase):
def test_ssl_ca_certs_without_certfile_raises(self):
with self.assertRaises(ValueError) as context:
ServerArgs(model_path="dummy", ssl_ca_certs="ca.pem")
self._validate_ssl(ssl_ca_certs="ca.pem")
self.assertIn("--ssl-ca-certs", str(context.exception))
def test_ssl_keyfile_password_without_certfile_raises(self):
with self.assertRaises(ValueError) as context:
ServerArgs(model_path="dummy", ssl_keyfile_password="secret")
self._validate_ssl(ssl_keyfile_password="secret")
self.assertIn("--ssl-keyfile-password", str(context.exception))
def test_ssl_keyfile_not_found_raises(self):
with self.assertRaises(ValueError) as context:
ServerArgs(
model_path="dummy",
self._validate_ssl(
ssl_keyfile="/nonexistent/key.pem",
ssl_certfile="/nonexistent/cert.pem",
)
@@ -817,8 +819,7 @@ class TestSSLArgs(unittest.TestCase):
def test_ssl_certfile_not_found_raises(self):
with tempfile.NamedTemporaryFile(suffix=".pem") as keyfile:
with self.assertRaises(ValueError) as context:
ServerArgs(
model_path="dummy",
self._validate_ssl(
ssl_keyfile=keyfile.name,
ssl_certfile="/nonexistent/cert.pem",
)
@@ -828,8 +829,7 @@ class TestSSLArgs(unittest.TestCase):
with tempfile.NamedTemporaryFile(suffix=".pem") as keyfile:
with tempfile.NamedTemporaryFile(suffix=".pem") as certfile:
with self.assertRaises(ValueError) as context:
ServerArgs(
model_path="dummy",
self._validate_ssl(
ssl_keyfile=keyfile.name,
ssl_certfile=certfile.name,
ssl_ca_certs="/nonexistent/ca.pem",
@@ -844,14 +844,13 @@ class TestSSLArgs(unittest.TestCase):
def test_enable_ssl_refresh_without_ssl_raises(self):
with self.assertRaises(ValueError) as context:
ServerArgs(model_path="dummy", enable_ssl_refresh=True)
self._validate_ssl(enable_ssl_refresh=True)
self.assertIn("--enable-ssl-refresh", str(context.exception))
self.assertIn("--ssl-certfile", str(context.exception))
@patch("os.path.isfile", return_value=True)
def test_enable_ssl_refresh_with_ssl_accepted(self, _mock_isfile):
server_args = ServerArgs(
model_path="dummy",
server_args = self._validate_ssl(
ssl_keyfile="key.pem",
ssl_certfile="cert.pem",
enable_ssl_refresh=True,
@@ -1160,8 +1159,7 @@ class TestPrefillOnlyDisableKvCache(unittest.TestCase):
- no context-parallel attention (CP writes to the pool via set_kv_buffer),
- no HiSparse (uses a different pool family),
- kv_cache_dtype != fp4_e2m1 (FP4 pool is a separate allocation path).
All other configurations must be rejected at __post_init__ time so users
get a clear error before model load.
All other configurations must be rejected before model load.
"""
def _base_kwargs(self, **overrides):
@@ -1175,47 +1173,54 @@ class TestPrefillOnlyDisableKvCache(unittest.TestCase):
kwargs.update(overrides)
return kwargs
def _validate_prefill_only_args(self, **overrides):
sa = ServerArgs(**self._base_kwargs(**overrides))
sa._handle_legacy_cp_arguments()
sa._validate_prefill_only_disable_kv_cache_args()
return sa
def test_valid_minimal_config_constructs(self):
sa = ServerArgs(**self._base_kwargs())
sa = self._validate_prefill_only_args()
self.assertTrue(sa.prefill_only_disable_kv_cache)
def test_rejects_when_not_embedding(self):
with self.assertRaisesRegex(ValueError, "requires --is-embedding"):
ServerArgs(**self._base_kwargs(is_embedding=False))
self._validate_prefill_only_args(is_embedding=False)
def test_rejects_when_chunked_prefill_size_not_minus_one(self):
with self.assertRaisesRegex(ValueError, "--chunked-prefill-size=-1"):
ServerArgs(**self._base_kwargs(chunked_prefill_size=8192))
self._validate_prefill_only_args(chunked_prefill_size=8192)
def test_rejects_when_radix_cache_enabled(self):
with self.assertRaisesRegex(ValueError, "--disable-radix-cache"):
ServerArgs(**self._base_kwargs(disable_radix_cache=False))
self._validate_prefill_only_args(disable_radix_cache=False)
def test_rejects_attn_cp_size_greater_than_one(self):
with self.assertRaisesRegex(ValueError, "--attn-cp-size"):
ServerArgs(**self._base_kwargs(attn_cp_size=2, tp_size=2))
self._validate_prefill_only_args(attn_cp_size=2, tp_size=2)
def test_rejects_prefill_context_parallel(self):
with self.assertRaisesRegex(ValueError, "--enable-prefill-cp"):
ServerArgs(**self._base_kwargs(enable_prefill_context_parallel=True))
self._validate_prefill_only_args(enable_prefill_context_parallel=True)
def test_rejects_hisparse(self):
with self.assertRaisesRegex(ValueError, "--enable-hisparse"):
ServerArgs(**self._base_kwargs(enable_hisparse=True))
self._validate_prefill_only_args(enable_hisparse=True)
def test_rejects_fp4_kv_cache(self):
with self.assertRaisesRegex(ValueError, "fp4_e2m1"):
ServerArgs(**self._base_kwargs(kv_cache_dtype="fp4_e2m1"))
self._validate_prefill_only_args(kv_cache_dtype="fp4_e2m1")
class TestSessionRadixCacheServerArgs(unittest.TestCase):
def test_requires_priority_radix_eviction_policy(self):
server_args = ServerArgs(
model_path="dummy",
enable_session_radix_cache=True,
radix_eviction_policy="lru",
)
with self.assertRaisesRegex(ValueError, "--radix-eviction-policy priority"):
ServerArgs(
model_path="dummy",
enable_session_radix_cache=True,
radix_eviction_policy="lru",
)
server_args._handle_cache_compatibility()
class TestCudaGraphConfigDataclassAccess(CustomTestCase):