Prhub

#47494 [Rust Frontend] Align sampling validation with Python

原始 PR 作者 reidliu41 合并时间 2026-07-28 20:02 文件变更 6 提交数 4 评论 10 代码增减 +260 / -2

执行摘要

对齐 Rust 前端采样参数验证至 Python 规范

Rust 前端未对采样参数做前面的校验,无效值(如 temperature=5.0、top_p=0.0)会直达引擎解码层并失败,而不是返回前端 400 Bad Request。详见 PR body 描述。

值得精读,尤其是 lower_sampling_params 中验证时机的选择以及错误类型与 API 响应格式的映射方式,为后续其他参数校验提供了可复用的模式。

讨论亮点

BugenZhao 在 review 中提出两个关键建议:

  • 将验证时机后移至 EngineCoreSamplingParams 构造完成后,而不是将各个参数独立传入,以避免接口膨胀和调用方遗漏;作者随后采纳。
  • 建议将 sampling.rs 放置在 rust/src/text/src/lower 目录下,与 logprobstoken_ids 等子模块平级,作者同样照做。
    讨论中 cinnamonica02 指出 gRPC 的 build_sampling_params 尚未被覆盖,作者回应已通过将校验下沉到 text 层使所有入口自动受益。

实现拆解

  1. 新增验证模块:在 rust/src/text/src/lower/sampling.rs 中定义 SamplingParamsError 枚举和七个独立的验证函数(validate_temperaturevalidate_top_pvalidate_min_pvalidate_frequency_penaltyvalidate_presence_penaltyvalidate_repetition_penalty,以及工具函数 validate_finitevalidate_closed_range),并提供公有函数 validate_resolved_sampling_params 统一验证已构造的 EngineCoreSamplingParams
  2. 集成到 lower 流程:在 rust/src/text/src/lower.rslower_sampling_params 函数中,于 EngineCoreSamplingParams 构造完成后之后插入 validate_resolved_sampling_params(&params)?; 的调用,这样 HTTP、gRPC 和原始 generate 路径都能共享同一份校验逻辑。
  3. 错误类型串联:在 rust/src/text/src/error.rs 中添加 SamplingParams 变体,在 rust/src/text/src/lib.rs 中重新导出 SamplingParamsError,使错误能向上层传播。
  4. 服务层错误映射:在 rust/src/server/src/error.rs 的测试中添加 sampling_params_validation_maps_to_invalid_request,确保 SamplingParamsError 正确映射为 400 BAD_REQUESTinvalid_request_error 类型。
  5. gRPC 集成测试:在 rust/src/server/src/grpc/tests.rs 中新增 unary_generate_invalid_sampling_params_returns_invalid_argument,验证 gRPC 端点对无效 top_p=2.0 返回 InvalidArgument 并包含错误参数名。
文件 模块 状态 重要度
rust/src/text/src/lower/sampling.rs text 层 added 8.8
rust/src/text/src/lower.rs text 层 modified 8.3
rust/src/server/src/grpc/tests.rs gRPC 层 modified 6.97
rust/src/server/src/error.rs 错误映射 modified 6.43
rust/src/text/src/error.rs 错误类型 modified 4.83
rust/src/text/src/lib.rs 库入口 modified 4.54

关键符号

validate_resolved_sampling_params validate_temperature validate_top_p validate_min_p validate_frequency_penalty validate_presence_penalty validate_repetition_penalty validate_finite validate_closed_range lower_sampling_params_rejects_invalid_sampling_ranges lower_sampling_params_rejects_non_finite_sampling_values lower_sampling_params_accepts_python_compatible_repetition_penalty_above_two unary_generate_invalid_sampling_params_returns_invalid_argument sampling_params_validation_maps_to_invalid_request

关键源码片段

rust/src/text/src/lower.rs core-logic

验证被嵌入到降层流程中的关键位置,确保所有端点共享

// rust/src/text/src/lower.rs ( 相关改动部分 )// 在 lower_sampling_params 末尾,EngineCoreSamplingParams 构造完成后立即验证
pub fn lower_sampling_params(...) -> Result<EngineCoreSamplingParams> {
    // ... 前面构造 params ...
    let params = EngineCoreSamplingParams {
        temperature,
        top_p,
        min_p,
        frequency_penalty,
        presence_penalty,
        repetition_penalty,
        // ... 其他字段 ...
    };
    // [!] 新增:在 vocab 范围校验之前,先校验采样参数本身的合法性
    validate_resolved_sampling_params(&params)?;
    validate_vocab_range(&params, &sampling_limits)?;
    Ok(params)
}#[cfg(test)]
mod tests {
    // ...    #[test]
    fn lower_sampling_params_rejects_invalid_sampling_ranges() {
        let cases = [
            ("temperature", SamplingParams { temperature: Some(5.0), ..Default::default() }),
            ("top_p", SamplingParams { top_p: Some(0.0), ..Default::default() }),
            ("min_p", SamplingParams { min_p: Some(2.0), ..Default::default() }),
            ("repetition_penalty", SamplingParams { repetition_penalty: Some(0.0), ..Default::default() }),
            ("frequency_penalty", SamplingParams { frequency_penalty: Some(100.0), ..Default::default() }),
            ("presence_penalty", SamplingParams { presence_penalty: Some(100.0), ..Default::default() }),
        ];        for (expected_parameter, sampling_params) in cases {
            let error = lower_sampling_params_with_limits(sampling_params, sample_sampling_limits())
                .unwrap_err();
            assert!(
                matches!(error, Error::SamplingParams(SamplingParamsError::OutOfRange { parameter, .. }) if parameter == expected_parameter),
                "{expected_parameter} should be rejected"
            );
        }
    }    #[test]
    fn lower_sampling_params_rejects_non_finite_sampling_values() {
        // 验证 Infinite 和 NaN 均被拒绝
        for (expected_parameter, sampling_params) in [
            ("temperature", SamplingParams { temperature: Some(f32::INFINITY), ..Default::default() }),
            ("repetition_penalty", SamplingParams { repetition_penalty: Some(f32::NAN), ..Default::default() }),
        ] {
            let error = lower_sampling_params_with_limits(sampling_params, sample_sampling_limits())
                .unwrap_err();
            assert!(
                matches!(error, Error::SamplingParams(SamplingParamsError::NotFinite { parameter, .. }) if parameter == expected_parameter),
                "{expected_parameter} should reject non-finite values"
            );
        }
    }    #[test]
    fn lower_sampling_params_accepts_python_compatible_repetition_penalty_above_two() {
        // Python 对 repetition_penalty 无上限,2.5 应被接受
        let params = lower_sampling_params_with_limits(
            SamplingParams { repetition_penalty: Some(2.5), ..Default::default() },
            sample_sampling_limits(),
        ).expect("repetition_penalty > 2 should be accepted");
    }
}
rust/src/server/src/grpc/tests.rs core-logic

验证 gRPC 端点也能正确拒绝无效采样参数

// rust/src/server/src/grpc/tests.rs ( 新增测试 )#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn unary_generate_invalid_sampling_params_returns_invalid_argument() {
    let (mut client, server_task, _engine_task) = grpc_test_server(
        b"engine-grpc-invalid-sampling",
        default_stream_output_specs(),
    ).await;    // 发送 top_p = 2.0(超出 (0,1] 范围)的请求,期望返回 InvalidArgument
    let status = client
        .generate(pb::GenerateRequest {
            request_id: "test-invalid-sampling".to_string(),
            model: "test-model".to_string(),
            prompt: Some(pb::generate_request::Prompt::Text("hi".to_string())),
            sampling: Some(pb::RandomSampling {
                top_p: 2.0,
                ..Default::default()
            }),
            ..Default::default()
        })
        .await
        .expect_err("should fail when top_p is out of range");    assert_eq!(status.code(), tonic::Code::InvalidArgument);
    assert!(status.message().contains("top_p"));    server_task.abort();
}

评论区精华

验证时机:引擎参数构造前 vs 构造后 设计

BugenZhao 建议在 `EngineCoreSamplingParams` 构造完成后统一验证,而不是在函数入口对每个参数分别调用,以避免接口膨胀和未来遗漏。

结论:作者采纳并移动了验证调用位置。 · 已解决

验证位置:顶层 vs text/llm 层 设计

BugenZhao 提出长期考虑将验证下沉到 `text` 或 `llm` 层,使 HTTP、gRPC、原始 generate 共享。

结论:作者最终将模块放在 `rust/src/text/src/lower/sampling.rs`,使所有路径受益。 · 已解决

gRPC 端点的覆盖范围 测试

cinnamonica02 指出 gRPC 的 `build_sampling_params` 转换函数也没有范围检查,可能未被覆盖。

结论:作者解释新的验证位于更低层的 `lower_sampling_params`,gRPC 路径最终也会经过该函数,因此自动受益。 · 已解决

风险与影响

低风险。新增的验证逻辑完全后置于 Python 端的范围定义一致,并且有充分的单元测试(范围、非有限值、repetition_penalty 上限开放)和 gRPC 集成测试覆盖。唯一潜在风险是验证函数中使用 f32::is_finite() 的精度问题,但目前与 Python 行为对齐,无明显隐患。

  • 用户侧:之前导致引擎内部错误的无效采样参数现在会返回清晰的 400 错误消息,提升了 API 一致性和调试体验。
  • 系统侧:避免无效请求浪费引擎资源,因为校验在降层阶段提前拦截。
  • 团队侧:统一了 Rust 前端与 Python 前端的校验逻辑,降低后续 drift 风险;新的验证模块结构清晰,易于扩展其他参数。
新增验证路径 API 兼容性 测试覆盖完善

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论