[3/N][Sync sglang-miles] TITO Support (#23751)

Co-authored-by: Jiajun Li <48857426+guapisolo@users.noreply.github.com>
This commit is contained in:
Yuzhen Zhou
2026-06-03 21:45:33 -04:00
committed by GitHub
co-authored by Jiajun Li
parent 084c6a7e2a
commit e03dfa8182
7 changed files with 217 additions and 9 deletions
@@ -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")