[rust-renderer] Standalone preprocessing (#36718)
Signed-off-by: Sage Ahrac <sagiahrak@gmail.com> Co-authored-by: Shangming Cai <csmthu@gmail.com> Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com> Co-authored-by: Rain Jiang <96632942+rainj-me@users.noreply.github.com>
This commit is contained in:
co-authored by
Shangming Cai
Liangsheng Yin
Rain Jiang
parent
6880a47955
commit
7b1c2ed0a4
@@ -0,0 +1,429 @@
|
||||
//! Protocol preparation invariants shared by rendering and inference.
|
||||
|
||||
use super::protocol::{
|
||||
ChatCompletionRequest, CompletionRequest, lower_chat_request, lower_text_completion_request,
|
||||
lower_token_ids_completion_request,
|
||||
};
|
||||
use super::test_utils::renderer_config;
|
||||
use crate::SamplingDefaults;
|
||||
|
||||
#[test]
|
||||
fn chat_lowering_preserves_template_controls_and_metadata() {
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"rid": "chat-lowering",
|
||||
"chat_template_kwargs": {"enable_thinking": false},
|
||||
"continue_final_message": true,
|
||||
"top_k": 17,
|
||||
"min_p": 0.2,
|
||||
"min_tokens": 3,
|
||||
"stop_regex": "END[0-9]",
|
||||
"ignore_eos": true,
|
||||
"skip_special_tokens": false,
|
||||
"return_meta_info": false,
|
||||
"bootstrap_host": "prefill",
|
||||
"bootstrap_port": 8998,
|
||||
"bootstrap_room": 42
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(request.model, "model");
|
||||
assert_eq!(
|
||||
request
|
||||
.chat_template_kwargs
|
||||
.as_ref()
|
||||
.and_then(|args| args.get("enable_thinking")),
|
||||
Some(&serde_json::Value::Bool(false))
|
||||
);
|
||||
assert!(request.continue_final_message);
|
||||
assert_eq!(request.sampling_overrides.top_k, Some(17));
|
||||
assert_eq!(request.sampling_overrides.min_p, Some(0.2));
|
||||
assert_eq!(request.sampling_overrides.min_tokens, Some(3));
|
||||
assert_eq!(request.sampling_overrides.ignore_eos, Some(true));
|
||||
assert_eq!(request.sampling_overrides.skip_special_tokens, Some(false));
|
||||
assert_eq!(request.extensions.return_meta_info, Some(false));
|
||||
|
||||
let (response_id, request) = lower_chat_request(&renderer_config(), request).unwrap();
|
||||
|
||||
assert_eq!(response_id, "chat-lowering");
|
||||
assert_eq!(request.metadata.bootstrap_host.as_deref(), Some("prefill"));
|
||||
assert_eq!(request.metadata.bootstrap_port, Some(8998));
|
||||
assert_eq!(request.metadata.bootstrap_room, Some(42));
|
||||
assert_eq!(request.sampling_params.top_k, 17);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chat_lowering_rejects_return_meta_info_until_supported() {
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"return_meta_info": true
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
let error = match lower_chat_request(&renderer_config(), request) {
|
||||
Ok(_) => panic!("return_meta_info=true must not be silently ignored"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(error.to_string().contains("return_meta_info"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn completion_sampling_defaults_follow_request_model_terminal_priority() {
|
||||
let mut config = renderer_config();
|
||||
config.default_sampling_params = SamplingDefaults {
|
||||
temperature: Some(0.6),
|
||||
top_p: Some(0.9),
|
||||
top_k: Some(32),
|
||||
min_p: Some(0.1),
|
||||
repetition_penalty: Some(1.1),
|
||||
};
|
||||
let omitted: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": "hello"
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, requests) = lower_text_completion_request(&config, &omitted).unwrap();
|
||||
let sampling = &requests[0].options.sampling_params;
|
||||
assert_eq!(sampling.temperature, 0.6);
|
||||
assert_eq!(sampling.top_p, 0.9);
|
||||
assert_eq!(sampling.top_k, 32);
|
||||
assert_eq!(sampling.min_p, 0.1);
|
||||
assert_eq!(sampling.repetition_penalty, 1.1);
|
||||
|
||||
let explicit: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": "hello",
|
||||
"temperature": 0.2,
|
||||
"top_p": 0.5,
|
||||
"top_k": 17,
|
||||
"min_p": 0.2,
|
||||
"repetition_penalty": 1.2
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, requests) = lower_text_completion_request(&config, &explicit).unwrap();
|
||||
let sampling = &requests[0].options.sampling_params;
|
||||
assert!((sampling.temperature - 0.2).abs() < 1e-6);
|
||||
assert!((sampling.top_p - 0.5).abs() < 1e-6);
|
||||
assert_eq!(sampling.top_k, 17);
|
||||
assert_eq!(sampling.min_p, 0.2);
|
||||
assert_eq!(sampling.repetition_penalty, 1.2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unsupported_sglang_fields_are_rejected_instead_of_ignored() {
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"input_ids": [1, 2, 3],
|
||||
"task": "domain"
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
let error = lower_chat_request(&renderer_config(), request)
|
||||
.unwrap_err()
|
||||
.to_string();
|
||||
|
||||
assert_eq!(error, "unsupported request fields: input_ids, task");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chat_modalities_keep_the_typed_openai_contract() {
|
||||
for modalities in [serde_json::json!("text"), serde_json::json!(["vision"])] {
|
||||
let request = serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"modalities": modalities
|
||||
});
|
||||
assert!(serde_json::from_value::<ChatCompletionRequest>(request).is_err());
|
||||
}
|
||||
|
||||
let text_request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"modalities": ["text"]
|
||||
}))
|
||||
.unwrap();
|
||||
lower_chat_request(&renderer_config(), text_request).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reasoning_inputs_normalize_with_python_precedence() {
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"reasoning_effort": "high",
|
||||
"reasoning": {"effort": "none", "enabled": true},
|
||||
"chat_template_kwargs": {"thinking": true}
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, request) = lower_chat_request(&renderer_config(), request).unwrap();
|
||||
let args = request.chat_template_args.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
serde_json::to_value(request.reasoning_effort).unwrap(),
|
||||
serde_json::json!("none")
|
||||
);
|
||||
assert_eq!(args.get("thinking"), Some(&serde_json::json!(true)));
|
||||
assert_eq!(args.get("enable_thinking"), Some(&serde_json::json!(false)));
|
||||
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"reasoning_effort": "0.5"
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, request) = lower_chat_request(&renderer_config(), request).unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(request.reasoning_effort).unwrap(),
|
||||
serde_json::json!(0.5)
|
||||
);
|
||||
assert_eq!(
|
||||
request
|
||||
.chat_template_args
|
||||
.as_ref()
|
||||
.and_then(|args| args.get("thinking")),
|
||||
Some(&serde_json::json!(true))
|
||||
);
|
||||
|
||||
for invalid in [serde_json::json!(true), serde_json::json!(1.0)] {
|
||||
let request = serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"reasoning_effort": invalid
|
||||
});
|
||||
assert!(serde_json::from_value::<ChatCompletionRequest>(request).is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn text_completion_lowering_attaches_batched_metadata_in_prompt_major_order() {
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": ["one", "two"],
|
||||
"n": 2,
|
||||
"rid": ["prompt-a", "prompt-b"],
|
||||
"cache_salt": ["tenant-a", "tenant-b"],
|
||||
"extra_key": ["", "batch"],
|
||||
"bootstrap_host": ["prefill-a", "prefill-b"],
|
||||
"bootstrap_port": [8998, null],
|
||||
"bootstrap_room": [41, 52],
|
||||
"priority": 7,
|
||||
"routed_dp_rank": 2
|
||||
}))
|
||||
.unwrap();
|
||||
let (response_id, requests) =
|
||||
lower_text_completion_request(&renderer_config(), &request).unwrap();
|
||||
|
||||
assert_eq!(response_id, "prompt-a");
|
||||
assert_eq!(
|
||||
requests
|
||||
.iter()
|
||||
.flat_map(|request| request.requests.iter())
|
||||
.map(|request| request.rid.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["prompt-a-0", "prompt-a-1", "prompt-b-0", "prompt-b-1"]
|
||||
);
|
||||
assert_eq!(
|
||||
requests[0].requests[0].metadata.cache_salt.as_deref(),
|
||||
Some("tenant-a")
|
||||
);
|
||||
assert_eq!(requests[0].requests[1].metadata.extra_key, None);
|
||||
assert_eq!(
|
||||
requests[1].requests[0].metadata.extra_key.as_deref(),
|
||||
Some("batch")
|
||||
);
|
||||
assert_eq!(requests[0].requests[0].metadata.bootstrap_port, Some(8998));
|
||||
assert_eq!(requests[1].requests[0].metadata.bootstrap_port, None);
|
||||
assert_eq!(requests[0].requests[1].metadata.bootstrap_room, Some(41));
|
||||
assert_eq!(requests[1].requests[1].metadata.bootstrap_room, Some(52));
|
||||
assert_eq!(requests[1].requests[1].metadata.routed_dp_rank, Some(2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn completion_lowering_validates_metadata_lengths_duplicates_and_scalar_rooms() {
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": ["one", "two"],
|
||||
"rid": ["duplicate", "duplicate"],
|
||||
"cache_salt": ["only-one"]
|
||||
}))
|
||||
.unwrap();
|
||||
let error = lower_text_completion_request(&renderer_config(), &request).unwrap_err();
|
||||
assert!(error.to_string().contains("duplicate request ID"));
|
||||
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": ["one", "two"],
|
||||
"cache_salt": ["only-one"]
|
||||
}))
|
||||
.unwrap();
|
||||
let error = lower_text_completion_request(&renderer_config(), &request).unwrap_err();
|
||||
assert!(error.to_string().contains("prompt batch size (2)"));
|
||||
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": ["one", "two"],
|
||||
"n": 2,
|
||||
"bootstrap_room": 90
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, requests) = lower_text_completion_request(&renderer_config(), &request).unwrap();
|
||||
assert_eq!(
|
||||
requests
|
||||
.iter()
|
||||
.flat_map(|request| request.requests.iter())
|
||||
.map(|request| request.metadata.bootstrap_room)
|
||||
.collect::<Vec<_>>(),
|
||||
[Some(90), Some(90), Some(91), Some(91)]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn completion_lowering_rejects_zero_max_tokens() {
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": "hello",
|
||||
"max_tokens": 0
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
let error = lower_text_completion_request(&renderer_config(), &request).unwrap_err();
|
||||
|
||||
assert_eq!(error.to_string(), "max_tokens must be positive");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn token_id_completion_lowering_attaches_batched_metadata() {
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": [[1, 2], [3]],
|
||||
"n": 2,
|
||||
"rid": ["tokens-a", "tokens-b"],
|
||||
"bootstrap_host": ["prefill-a", "prefill-b"],
|
||||
"bootstrap_port": [8998, 8999],
|
||||
"bootstrap_room": [41, 52]
|
||||
}))
|
||||
.unwrap();
|
||||
let (response_id, requests) =
|
||||
lower_token_ids_completion_request(&renderer_config(), &request).unwrap();
|
||||
|
||||
assert_eq!(response_id, "tokens-a");
|
||||
assert_eq!(requests[2].rid, "tokens-b-0");
|
||||
assert_eq!(requests[2].input_ids, [3]);
|
||||
assert_eq!(
|
||||
requests[2].metadata.bootstrap_host.as_deref(),
|
||||
Some("prefill-b")
|
||||
);
|
||||
assert_eq!(requests[2].metadata.bootstrap_port, Some(8999));
|
||||
assert_eq!(requests[3].metadata.bootstrap_room, Some(52));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn route_operations_decode_tokens_without_http() {
|
||||
use super::{OpenAIService, OperationResponse};
|
||||
use crate::engine::{
|
||||
GenerateTransport, GenerationService, TokenDecoder, TokenDelta, TokenStream,
|
||||
};
|
||||
use crate::{
|
||||
DynamoTokenizer, GenerateRequest, GenerationFinishReason, RendererService, ResponseError,
|
||||
};
|
||||
use futures::{StreamExt, future::BoxFuture};
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
struct MemoryTransport(Mutex<Vec<GenerateRequest>>);
|
||||
impl GenerateTransport for MemoryTransport {
|
||||
fn generate(
|
||||
&self,
|
||||
request: GenerateRequest,
|
||||
) -> BoxFuture<'_, Result<TokenStream, ResponseError>> {
|
||||
Box::pin(async move {
|
||||
self.0.lock().unwrap().push(request);
|
||||
Ok(futures::stream::iter([Ok(TokenDelta {
|
||||
token_ids: vec![104],
|
||||
prompt_tokens: 5,
|
||||
completion_tokens: 1,
|
||||
finish_reason: Some(GenerationFinishReason::Length),
|
||||
..Default::default()
|
||||
})])
|
||||
.boxed())
|
||||
})
|
||||
}
|
||||
}
|
||||
async fn values<U: serde::Serialize, C: serde::Serialize>(
|
||||
result: OperationResponse<U, C>,
|
||||
) -> Vec<serde_json::Value> {
|
||||
match result {
|
||||
OperationResponse::Unary(value) => vec![serde_json::to_value(value).unwrap()],
|
||||
OperationResponse::Stream(stream) => {
|
||||
stream
|
||||
.map(|value| serde_json::to_value(value.unwrap()).unwrap())
|
||||
.collect()
|
||||
.await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let tokenizer = crate::engine::test_utils::tiny_tokenizer();
|
||||
let prompt_ids = tokenizer.encode("hello").unwrap().token_ids().to_vec();
|
||||
let transport = Arc::new(MemoryTransport(Mutex::new(Vec::new())));
|
||||
let renderer = Arc::new(RendererService::with_tokenizer(
|
||||
renderer_config(),
|
||||
Arc::new(DynamoTokenizer::new(tokenizer.clone(), tokenizer.clone())),
|
||||
1,
|
||||
1,
|
||||
));
|
||||
let service = OpenAIService::new(
|
||||
renderer,
|
||||
GenerationService::new(transport.clone(), TokenDecoder::new(tokenizer)),
|
||||
);
|
||||
for chat in [false, true] {
|
||||
for stream in [false, true] {
|
||||
let mut body =
|
||||
serde_json::json!({"model": "model", "n": 2, "max_tokens": 4, "stream": stream});
|
||||
let responses = if chat {
|
||||
body["messages"] = serde_json::json!([{"role": "user", "content": "hello"}]);
|
||||
values(
|
||||
service
|
||||
.chat(serde_json::from_value(body).unwrap())
|
||||
.await
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
body["prompt"] = serde_json::json!(prompt_ids);
|
||||
body["echo"] = serde_json::json!(true);
|
||||
values(
|
||||
service
|
||||
.complete(serde_json::from_value(body).unwrap())
|
||||
.await
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
};
|
||||
let mut texts = [String::new(), String::new()];
|
||||
let mut finished = [false; 2];
|
||||
for response in responses {
|
||||
for choice in response["choices"].as_array().unwrap() {
|
||||
let index = choice["index"].as_u64().unwrap() as usize;
|
||||
let text = if chat {
|
||||
&choice[if stream { "delta" } else { "message" }]["content"]
|
||||
} else {
|
||||
&choice["text"]
|
||||
};
|
||||
texts[index].push_str(text.as_str().unwrap_or_default());
|
||||
if let Some(reason) = choice["finish_reason"].as_str() {
|
||||
assert_eq!(reason, "length");
|
||||
finished[index] = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
assert_eq!(texts, [if chat { "h" } else { "helloh" }; 2]);
|
||||
assert_eq!(finished, [true; 2]);
|
||||
}
|
||||
}
|
||||
let requests = transport.0.lock().unwrap();
|
||||
assert_eq!(requests.len(), 8);
|
||||
assert!(requests.iter().all(|request| !request.input_ids.is_empty()));
|
||||
}
|
||||
Reference in New Issue
Block a user