[3/N][Sync sglang-miles] TITO Support (#23751)
Co-authored-by: Jiajun Li <48857426+guapisolo@users.noreply.github.com>
This commit is contained in:
@@ -179,6 +179,20 @@ class TestChatCompletionRequest(unittest.TestCase):
|
||||
self.assertFalse(request.stream_reasoning)
|
||||
self.assertEqual(request.chat_template_kwargs, {"custom_param": "value"})
|
||||
|
||||
def test_chat_completion_tito_extensions(self):
|
||||
"""Test chat completion with pre-tokenized prompt extensions."""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
request = ChatCompletionRequest(
|
||||
model="test-model",
|
||||
messages=messages,
|
||||
input_ids=[101, 102, 103],
|
||||
return_prompt_token_ids=True,
|
||||
return_meta_info=True,
|
||||
)
|
||||
self.assertEqual(request.input_ids, [101, 102, 103])
|
||||
self.assertTrue(request.return_prompt_token_ids)
|
||||
self.assertTrue(request.return_meta_info)
|
||||
|
||||
def test_chat_completion_reasoning_effort(self):
|
||||
"""Test chat completion with reasoning effort"""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
@@ -371,6 +385,28 @@ class TestModelSerialization(unittest.TestCase):
|
||||
self.assertIn("hidden_states", data["choices"][0])
|
||||
self.assertEqual(data["choices"][0]["hidden_states"], [0.1, 0.2, 0.3])
|
||||
|
||||
def test_prompt_token_ids_and_meta_info_serialization(self):
|
||||
"""Test that prompt_token_ids and meta_info serialize only when set."""
|
||||
default_choice = ChatCompletionResponseChoice(
|
||||
index=0,
|
||||
message=ChatMessage(role="assistant", content="Hello"),
|
||||
finish_reason="stop",
|
||||
)
|
||||
default_data = default_choice.model_dump()
|
||||
self.assertNotIn("prompt_token_ids", default_data)
|
||||
self.assertNotIn("meta_info", default_data)
|
||||
|
||||
choice = ChatCompletionResponseChoice(
|
||||
index=0,
|
||||
message=ChatMessage(role="assistant", content="Hello"),
|
||||
finish_reason="stop",
|
||||
prompt_token_ids=[1, 2, 3],
|
||||
meta_info={"prompt_tokens": 3},
|
||||
)
|
||||
data = choice.model_dump()
|
||||
self.assertEqual(data["prompt_token_ids"], [1, 2, 3])
|
||||
self.assertEqual(data["meta_info"], {"prompt_tokens": 3})
|
||||
|
||||
|
||||
class TestFunctionDeferLoading(unittest.TestCase):
|
||||
"""Test defer_loading field behavior on Function/Tool."""
|
||||
|
||||
@@ -147,6 +147,55 @@ class ServingChatTestCase(unittest.TestCase):
|
||||
self.assertFalse(adapted.stream)
|
||||
self.assertEqual(processed, self.basic_req)
|
||||
|
||||
def test_convert_to_internal_request_rejects_stream_return_prompt_token_ids(self):
|
||||
req = ChatCompletionRequest(
|
||||
model="x",
|
||||
messages=[{"role": "user", "content": "Hi?"}],
|
||||
stream=True,
|
||||
return_prompt_token_ids=True,
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
ValueError, "return_prompt_token_ids is not supported with streaming"
|
||||
):
|
||||
self.chat._convert_to_internal_request(req, self.fastapi_request)
|
||||
|
||||
def test_convert_to_internal_request_rejects_stream_return_meta_info(self):
|
||||
req = ChatCompletionRequest(
|
||||
model="x",
|
||||
messages=[{"role": "user", "content": "Hi?"}],
|
||||
stream=True,
|
||||
return_meta_info=True,
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
ValueError, "return_meta_info is not supported with streaming"
|
||||
):
|
||||
self.chat._convert_to_internal_request(req, self.fastapi_request)
|
||||
|
||||
def test_convert_to_internal_request_input_ids_bypasses_template(self):
|
||||
self.tm.tokenizer = None
|
||||
req = ChatCompletionRequest(
|
||||
model="x",
|
||||
messages=[{"role": "user", "content": "Hi?"}],
|
||||
input_ids=[101, 102, 103],
|
||||
stop=["STOP"],
|
||||
return_prompt_token_ids=True,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||||
) as conv_mock:
|
||||
adapted, processed = self.chat._convert_to_internal_request(
|
||||
req, self.fastapi_request
|
||||
)
|
||||
|
||||
self.assertEqual(processed, req)
|
||||
self.assertEqual(adapted.input_ids, [101, 102, 103])
|
||||
self.assertTrue(adapted.return_prompt_token_ids)
|
||||
self.assertEqual(adapted.sampling_params["stop"], ["STOP"])
|
||||
conv_mock.assert_not_called()
|
||||
|
||||
def test_kimi_tool_call_keeps_default_reasoning(self):
|
||||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||||
toggle_param="thinking", default_enabled=True
|
||||
@@ -1290,6 +1339,39 @@ class ServingChatTestCase(unittest.TestCase):
|
||||
},
|
||||
)
|
||||
|
||||
def test_non_streaming_chat_response_returns_requested_prompt_ids_and_meta_info(
|
||||
self,
|
||||
):
|
||||
req = ChatCompletionRequest(
|
||||
model="x",
|
||||
messages=[{"role": "user", "content": "Hi?"}],
|
||||
return_prompt_token_ids=True,
|
||||
return_meta_info=True,
|
||||
)
|
||||
ret = [
|
||||
{
|
||||
"text": "Answer",
|
||||
"prompt_token_ids": [11, 12, 13],
|
||||
"meta_info": {
|
||||
"id": "chatcmpl-token-ids",
|
||||
"prompt_tokens": 3,
|
||||
"completion_tokens": 1,
|
||||
"cached_tokens": 0,
|
||||
"finish_reason": {"type": "stop", "matched": None},
|
||||
"weight_version": "default",
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
response = self.chat._build_chat_response(req, ret, created=123)
|
||||
choice = response.choices[0]
|
||||
|
||||
self.assertEqual(choice.prompt_token_ids, [11, 12, 13])
|
||||
self.assertEqual(choice.meta_info, ret[0]["meta_info"])
|
||||
dumped_choice = response.model_dump()["choices"][0]
|
||||
self.assertEqual(dumped_choice["prompt_token_ids"], [11, 12, 13])
|
||||
self.assertEqual(dumped_choice["meta_info"], ret[0]["meta_info"])
|
||||
|
||||
def test_streaming_cached_tokens_details_emits_sglext(self):
|
||||
"""Test that streaming chat responses emit cached token details in sglext."""
|
||||
|
||||
@@ -1709,7 +1791,7 @@ class ServingChatTestCase(unittest.TestCase):
|
||||
|
||||
response = self.chat._build_chat_response(req, [ret_item], created=0)
|
||||
msg = response.choices[0].message
|
||||
self.assertIsNone(msg.content)
|
||||
self.assertEqual(msg.content, "")
|
||||
self.assertEqual(msg.reasoning_content, "42")
|
||||
|
||||
# --- poolside_v1 (Laguna-XS.2) regression tests ---
|
||||
|
||||
@@ -539,6 +539,19 @@ class TestGenerateReqInputNormalization(CustomTestCase):
|
||||
self.assertEqual(item0.custom_logit_processor, "processor1")
|
||||
self.assertEqual(item0.return_hidden_states, True)
|
||||
|
||||
def test_getitem_preserves_return_prompt_token_ids(self):
|
||||
"""Batch subrequests must keep the prompt-token-id return flag."""
|
||||
req = GenerateReqInput(
|
||||
input_ids=[[1, 2, 3], [4, 5, 6]],
|
||||
sampling_params=[{}, {}],
|
||||
rid=["id1", "id2"],
|
||||
return_prompt_token_ids=True,
|
||||
)
|
||||
req.normalize_batch_and_arguments()
|
||||
|
||||
self.assertTrue(req[0].return_prompt_token_ids)
|
||||
self.assertTrue(req[1].return_prompt_token_ids)
|
||||
|
||||
def test_regenerate_rid(self):
|
||||
"""Test the regenerate_rid method."""
|
||||
req = GenerateReqInput(text="Hello")
|
||||
|
||||
Reference in New Issue
Block a user