Prhub

#46839 [Bugfix][Rust Frontend] Reject prompt_logprobs for streaming generate

原始 PR 作者 reidliu41 合并时间 2026-06-30 13:10 文件变更 2 提交数 3 评论 3 代码增减 +103 / -8

执行摘要

拒绝流式 generate 请求中的 prompt_logprobs 参数

/inference/v1/generate 中,当请求同时设置 stream=truesampling_params.prompt_logprobs > 0=-1 时,后端返回 200 并开始 SSE 流,但流式响应的 shape 不包含 prompt_logprobs 字段,导致该参数被静默忽略。本次变更拒绝该不支持的参数组合,让客户端立即获知错误。

建议所有使用 Rust 前端的用户更新,该修复符合 API 设计一致性。值得关注的点:作者通过将 if-let 与链转换为嵌套 if,实现了更清晰的逻辑扩展,便于后续添加类似的互斥参数检查。

讨论亮点

在 review 中,chatgpt-codex-connector[bot] 指出最初的变更未拒绝 prompt_logprobs=0,而 0 也是显式请求 prompt_logprobs,同样不应支持流式。作者 reidliu41 确认并修复,最终版本拒绝所有 Some(prompt_logprobs) 的流式请求。

实现拆解

  1. rust/src/server/src/routes/inference/generate/validate.rsvalidate_request_compat 函数中,将原有的 if let 与链转换成嵌套 if,并新增分支:当 request.streamtrueprompt_logprobs 被设置时,返回 400 Bad Request
  2. 保持对 prompt_logprobs 值合法性的检查(非负或 -1)不变。
  3. 在同一个文件中添加两个单元测试:validate_request_compat_rejects_streaming_prompt_logprobs 验证流式请求被拒绝(含 01-1 三种情况),validate_request_compat_accepts_non_stream_prompt_logprobs 验证非流式请求仍然通过。
  4. rust/src/server/src/routes/tests.rs 中添加集成测试 raw_generate_rejects_streaming_prompt_logprobs,模拟 HTTP 请求并断言返回状态码和错误信息。
文件 模块 状态 重要度
rust/src/server/src/routes/inference/generate/validate.rs 路由验证 modified 7.19
rust/src/server/src/routes/tests.rs 集成测试 modified 6.35

关键符号

validate_request_compat validate_request_compat_rejects_streaming_prompt_logprobs validate_request_compat_accepts_non_stream_prompt_logprobs raw_generate_rejects_streaming_prompt_logprobs

关键源码片段

rust/src/server/src/routes/inference/generate/validate.rs core-logic

核心验证逻辑,新增针对流式与 prompt_logprobs 冲突的拒绝检查,并调整条件判断结构。

// 验证生成请求的兼容性
pub(super) 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 只能在 stream=true 时设置
    if request.stream_options.is_some() && !request.stream {
        bail_invalid_request!(
            param = "stream_options",
            "stream_options are only supported when stream=true."
        );
    }    // token_ids 不能为空
    if request.token_ids.is_empty() {
        bail_invalid_request!(
            param = "token_ids",
            "token_ids must contain at least one token ID."
        );
    }    // max_tokens 必须大于 0
    if request.sampling_params.max_tokens == Some(0) {
        bail_invalid_request!(
            param = "sampling_params",
            "max_tokens must be greater than 0."
        );
    }    // 当 prompt_logprobs 被显式设置时:
    if let Some(prompt_logprobs) = request.sampling_params.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."
            );
        }        // 拒绝流式生成时设置 prompt_logprobs,因为流式输出不包含 prompt_logprobs 字段
        if request.stream {
            bail_invalid_request!(
                param = "sampling_params",
                "`prompt_logprobs` are not available when `stream=true`."
            );
        }
    }    Ok(())
}
rust/src/server/src/routes/tests.rs test-coverage

集成测试,验证 HTTP 路由层对非法请求返回 400 状态,确保端到端正确性。

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn raw_generate_rejects_streaming_prompt_logprobs() {
    let mut app = test_app().await;    // 测试 prompt_logprobs = 0 和 =1 两种情况( -1 已在单元测试中验证)
    for prompt_logprobs in [0, 1] {
        let response = app
            .call(
                Request::builder()
                    .method("POST")
                    .uri("/inference/v1/generate")
                    .header("content-type", "application/json")
                    .body(Body::from(
                        json!({
                            "model": "Qwen/Qwen1.5-0.5B-Chat",
                            "token_ids": [11, 22],
                            "stream": true,
                            "sampling_params": {
                                "prompt_logprobs": prompt_logprobs
                            }
                        })
                        .to_string(),
                    ))
                    .expect("build request"),
            )
            .await
            .expect("call app");        // 验证返回 400 错误以及错误信息
        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
        let body = to_bytes(response.into_body(), usize::MAX).await.expect("read body");
        let json: serde_json::Value = serde_json::from_slice(&body).expect("decode json");
        assert_eq!(json["error"]["param"], "sampling_params");
        assert_eq!(
            json["error"]["message"],
            "`prompt_logprobs` are not available when `stream=true`."
        );
    }
}

评论区精华

拒绝 prompt_logprobs=0 的流式请求 正确性

chatgpt-codex-connector[bot] 指出最初的变更未拒绝 prompt_logprobs=0,而 0 也是显式请求 prompt_logprobs,同样不应支持流式。

结论:作者 reidliu41 确认并修复,最终版本拒绝所有 Some(prompt_logprobs) 的流式请求。 · 已解决

风险与影响

该变更很小,风险低。主要风险是破坏向后兼容性:之前使用流式且传入 prompt_logprobs 的客户端将从静默忽略变为 400 错误,但这是预期行为变更,且更早暴露错误。无性能影响。

影响范围限于 Rust 前端的单一路径 /inference/v1/generate。影响程度小,变更清晰,测试覆盖充分。对系统其他部分无影响。

API 行为变更 流式兼容性

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论