Prhub

#52844 [Bugfix][Rust Frontend] Reject n > 1 in the `/inference/v1/generate` route

原始 PR 作者 qgallouedec 合并时间 2026-08-20 15:41 文件变更 6 提交数 2 评论 6 代码增减 +85 / -9

执行摘要

Rust 前端 generate 路由显式拒绝 n>1,消除静默丢结果

Issue #52843 报告:VLLM_USE_RUST_FRONTEND=1 时,非流式请求 /inference/v1/generate 携带 sampling_params: {"n": 4} 返回 HTTP 200 但只有单条 choices(index: 0),“No error/warning but the other completions are silently missing”。根因是共享的 vllm_text::SamplingParams 故意不含 n 字段(并行采样在 v1 中由 Python 前端 ParentRequest fan-out 实现),反序列化时该键被丢弃,而 collect_generate 硬编码返回单条选择。Rust OpenAI 路由(completions/chat_completions 的 validate.rs)早已显式拒绝 n > 1,raw generate 路由应保持一致。本 PR 是 Python 前端修复 #52399 的 Rust 对应实现。

值得精读的小型修复。它展示了在不动共享 crate 的前提下,用 serde(flatten) 包装类型在路由层捕获共享类型故意省略字段的通用模式,对参与 vllm-text Rust 前端开发的工程师有参考价值。值得关注的决策点:unwrap_or(1) 的默认值语义、与 OpenAI 路由校验文案的一致性、render 复用路径显式 n: None 的构造方式,以及评审中 > 1 vs != 1 的边界收紧。

讨论亮点

唯一实质性 review 评论来自 BugenZhao:在 validate.rs 的 diff 上给出 suggestion,将条件从 request.sampling_params.n.unwrap_or(1) > 1 改为 request.sampling_params.n.unwrap_or(1) != 1。理由是 > 1 会放行 n == 0 这类同样非法的值,而 != 1 与 OpenAI 路由 Only n=1 is supported. 的语义完全一致。该建议已在第二个 commit(BugenZhao 提交的 validate.rs 更新)中落实。Codex 自动 review 结论为 “Didn't find any major issues”,BugenZhao 最终 APPROVED 并留言 “LGTM, thanks!”。

实现拆解

实现分四步:

  1. 数据契约(types.rs):新增 GenerateSamplingParams { n: Option<u32>, #[serde(flatten)] inner: SamplingParams },并将 GenerateRequest.sampling_params 字段类型由 SamplingParams 改为 GenerateSamplingParamsserde(flatten) 保证 JSON 结构不变,仅把 n 单独捕获出来,为校验提供数据入口,且不动共享 crate。
  2. 校验逻辑(validate.rs):在 validate_request_compat 中于 stream_options 校验后、token_ids 校验前插入 if request.sampling_params.n.unwrap_or(1) != 1 { bail_invalid_request!(param = "n", "Only n=1 is supported.") };同时把 max_tokensprompt_logprobs 的访问改为 sampling_params.inner.*unwrap_or(1) 保证未传 n 时兼容旧行为。评审中 BugenZhao 建议用 != 1 而非 > 1,以覆盖 n == 0 等非法值,已采纳。
  3. 消费端适配(convert.rs、generate.rs、render.rs)prepare_generate_requestrequest.sampling_params.inner 读取 logprobs/prompt_logprobs 并透传 inner 给引擎;lower_render_request 构造请求时显式 GenerateSamplingParams { n: None, inner: text_request.sampling_params }generate.rs 模块导出新增 GenerateSamplingParams,保证 render 内部复用路径行为不变。
  4. 测试与验证:validate.rs 内联两个单元测试(n=4 拒绝、n=1 接受),tests.rs 新增 HTTP 路由级测试 raw_generate_rejects_parallel_sampling(断言 400 且 error.param == "n")。PR 说明 cargo test -p vllm-server --lib、fmt、clippy 全绿,Buildkite CI #84762 已触发。无配置与部署配套改动。
文件 模块 状态 重要度
rust/src/server/src/routes/inference/generate/validate.rs 请求路由 modified 7.18
rust/src/server/src/routes/inference/generate/types.rs 请求路由 modified 5.78
rust/src/server/src/routes/tests.rs 请求路由 modified 6.27
rust/src/server/src/routes/render.rs 请求路由 modified 4.83
rust/src/server/src/routes/inference/generate/convert.rs 请求路由 modified 4.62
rust/src/server/src/routes/inference/generate.rs 请求路由 modified 3.92

关键符号

validate_request_compat validate_request_compat_rejects_parallel_sampling validate_request_compat_accepts_explicit_n_one raw_generate_rejects_parallel_sampling prepare_generate_request lower_render_request

关键源码片段

rust/src/server/src/routes/inference/generate/validate.rs entrypoint

校验主路径:新增对 sampling_params.n 的显式校验(拒绝 n != 1),并把 max_tokens、prompt_logprobs 访问迁移到 inner;同时内联两个单元测试覆盖 n=4 拒绝与 n=1 接受。

// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM projectuse super::types::GenerateRequest;
use crate::error::{ApiError, bail_invalid_request};/// 为 Rust token generate 路由强制执行最小兼容性契约。
pub(crate) fn validate_request_compat(
    request: &GenerateRequest,
    served_model_names: &[String],
) -> Result<(), ApiError> {
    // 模型名校验:不在已服务模型列表内则拒绝。
    if let Some(model) = request.model.as_ref()
        && !served_model_names.iter().any(|n| n == model)
    {
        return Err(ApiError::model_not_found(model.clone()));
    }    // stream_options 仅在流式请求下有意义。
    if request.stream_options.is_some() && !request.stream {
        bail_invalid_request!(
            param = "stream_options",
            "stream_options are only supported when stream=true."
        );
    }    // Rust 前端不实现并行采样(v1 中由 Python 前端 ParentRequest fan-out 承担),
    // `vllm_text::SamplingParams` 刻意省略 `n` 字段,因此这里显式捕获并拒绝 `n != 1`,
    // 避免反序列化时静默丢键、最终只返回单条结果。`unwrap_or(1)` 保证未传 `n` 时不误伤。
    if request.sampling_params.n.unwrap_or(1) != 1 {
        bail_invalid_request!(param = "n", "Only n=1 is supported.");
    }    if request.token_ids.is_empty() {
        bail_invalid_request!(
            param = "token_ids",
            "token_ids must contain at least one token ID."
        );
    }    // 其余采样参数统一通过包装类型的内层 SamplingParams 访问。
    if request.sampling_params.inner.max_tokens == Some(0) {
        bail_invalid_request!(
            param = "sampling_params",
            "max_tokens must be greater than 0."
        );
    }    if let Some(prompt_logprobs) = request.sampling_params.inner.prompt_logprobs {
        // prompt_logprobs 需为非负值或 -1,且流式请求不支持。
        if prompt_logprobs < 0 && prompt_logprobs != -1 {
            bail_invalid_request!(
                param = "sampling_params",
                "`prompt_logprobs` must be a non-negative value or -1."
            );
        }        if request.stream {
            bail_invalid_request!(
                param = "sampling_params",
                "`prompt_logprobs` are not available when `stream=true`."
            );
        }
    }    Ok(())
}
rust/src/server/src/routes/inference/generate/types.rs data-contract

数据契约核心:新增 GenerateSamplingParams 包装类型,用 serde(flatten) 在路由层捕获共享 vllm_text::SamplingParams 刻意省略的 n 字段,是本修复的技术基础。

/// 供 generate API(token-in/token-out)使用的采样参数类型。
///
/// 包装 [`SamplingParams`] 以额外捕获 `n` 字段:共享的 `vllm_text::SamplingParams`
/// 故意不包含 `n`(并行采样由更高层处理,Rust 前端不实现)。若路由层不捕获 `n`,
/// 反序列化时该键会被静默丢弃,最终只返回单条生成结果。
#[serde_with::skip_serializing_none]
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct GenerateSamplingParams {
    /// 输出序列数量,本路由仅支持 1。
    pub n: Option<u32>,
    /// 其余采样参数经 flatten 后与原 JSON 结构保持一致,直接透传给引擎。
    #[serde(flatten)]
    pub inner: SamplingParams,
}// GenerateRequest 中的使用方式:sampling_params 字段类型改为 GenerateSamplingParams,
// 反序列化时 `n` 被捕获到外层,其余键全部下沉到 inner,校验与引擎透传互不干扰。
pub struct GenerateRequest {
    // ...
    pub sampling_params: GenerateSamplingParams,
    // ...
}

评论区精华

n 校验边界:> 1 还是 != 1 正确性

BugenZhao 在 validate.rs 的 diff 上给出 suggestion,将条件从 `request.sampling_params.n.unwrap_or(1) > 1` 改为 `request.sampling_params.n.unwrap_or(1) != 1`。原写法会放行 `n == 0` 等非法值,与 OpenAI 路由 `Only n=1 is supported.` 的语义不完全一致。

结论:已采纳。第二个 commit(BugenZhao 提交 validate.rs 更新)落实 `!= 1`,拒绝一切非 1 的 `n`,使校验与 OpenAI completions/chat_completions 路由完全对齐。 · 已解决

Codex 自动审查与维护者批准 other

BugenZhao 触发 `@codex review`,Codex 结论为未发现重大问题;随后 BugenZhao 手动 APPROVED 并留言 LGTM。

结论:无遗留问题,PR 由 BugenZhao 合并。 · 已解决

风险与影响

  1. API 行为变更n > 1 从静默返回 200 单条结果变为显式 400,依赖旧行为的客户端会开始收到错误;这是修复本意,且与 OpenAI 路由一致,属于合理收紧。
  2. 类型包装连带影响GenerateRequest.sampling_params 类型变化波及所有内部访问点(validate.rs、convert.rs、render.rs 已同步),未来新增消费代码若直接访问 sampling_params.max_tokens 等字段,会在编译期报错(字段已移到 inner),Rust 类型系统兜底,运行期无隐患。
  3. 测试内联风险:测试全部内联在 validate.rs 的 mod tests 与 routes/tests.rs,未新增独立测试文件,但与现有测试风格一致,覆盖了协议层和边界值,回归风险可控。

用户:仅影响 VLLM_USE_RUST_FRONTEND=1 下调用 raw generate 传 n > 1 的客户端,从静默丢结果变为可诊断的 400 错误(error.param = "n")。系统:改动限定在 Rust 前端路由层(generate.rs 模块与 render.rs),引擎侧与共享 vllm_text crate 行为不变,OpenAI 兼容路由无影响。团队:为 Rust 前端补齐与 Python 前端一致的参数校验契约,并保留 n 字段的显式占位,为将来在 Rust 前端实现并行采样(ParentRequest fan-out)预留数据入口。

API 行为变更:n>1 从静默 200 变为显式 400 采样参数迁移到 inner 后新增消费点需同步适配

关联 Issue

#52843 [Rust frontend] /inference/v1/generate silently ignores n > 1; should reject it

完整报告

参与讨论