| import unittest |
|
|
| import openai |
|
|
| from sglang.srt.utils import kill_process_tree |
| from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci |
| from sglang.test.test_utils import ( |
| DEFAULT_SMALL_MODEL_NAME_FOR_TEST, |
| DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, |
| DEFAULT_URL_FOR_TEST, |
| CustomTestCase, |
| popen_launch_server, |
| ) |
|
|
| register_cuda_ci(est_time=38, suite="stage-b-test-large-1-gpu") |
| register_amd_ci(est_time=31, suite="stage-b-test-small-1-gpu-amd") |
|
|
|
|
| class TestRequestLengthValidation(CustomTestCase): |
| @classmethod |
| def setUpClass(cls): |
| cls.base_url = DEFAULT_URL_FOR_TEST |
| cls.api_key = "sk-123456" |
|
|
| |
| cls.process = popen_launch_server( |
| DEFAULT_SMALL_MODEL_NAME_FOR_TEST, |
| cls.base_url, |
| timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, |
| api_key=cls.api_key, |
| other_args=("--max-total-tokens", "1000", "--context-length", "1000"), |
| ) |
|
|
| @classmethod |
| def tearDownClass(cls): |
| kill_process_tree(cls.process.pid) |
|
|
| def test_input_length_longer_than_context_length(self): |
| client = openai.Client(api_key=self.api_key, base_url=f"{self.base_url}/v1") |
|
|
| long_text = "hello " * 1200 |
|
|
| with self.assertRaises(openai.BadRequestError) as cm: |
| client.chat.completions.create( |
| model=DEFAULT_SMALL_MODEL_NAME_FOR_TEST, |
| messages=[ |
| {"role": "user", "content": long_text}, |
| ], |
| temperature=0, |
| ) |
|
|
| self.assertIn("is longer than the model's context length", str(cm.exception)) |
|
|
| def test_input_length_longer_than_maximum_allowed_length(self): |
| client = openai.Client(api_key=self.api_key, base_url=f"{self.base_url}/v1") |
|
|
| long_text = "hello " * 999 |
|
|
| with self.assertRaises(openai.BadRequestError) as cm: |
| client.chat.completions.create( |
| model=DEFAULT_SMALL_MODEL_NAME_FOR_TEST, |
| messages=[ |
| {"role": "user", "content": long_text}, |
| ], |
| temperature=0, |
| ) |
|
|
| self.assertIn("is longer than the model's context length", str(cm.exception)) |
|
|
| def test_max_tokens_validation(self): |
| client = openai.Client(api_key=self.api_key, base_url=f"{self.base_url}/v1") |
|
|
| long_text = "hello " |
|
|
| with self.assertRaises(openai.BadRequestError) as cm: |
| client.chat.completions.create( |
| model=DEFAULT_SMALL_MODEL_NAME_FOR_TEST, |
| messages=[ |
| {"role": "user", "content": long_text}, |
| ], |
| temperature=0, |
| max_tokens=1200, |
| ) |
|
|
| self.assertIn( |
| "max_completion_tokens is too large", |
| str(cm.exception), |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|