Skip to main content

openai_protocol/
sampling_params.rs

1use serde::{Deserialize, Serialize};
2use serde_json::{Map, Value};
3use validator::Validate;
4
5use super::common::StringOrArray;
6
7/// Sampling parameters for text generation
8#[serde_with::skip_serializing_none]
9#[derive(Debug, Clone, Deserialize, Serialize, Default, Validate, schemars::JsonSchema)]
10#[validate(schema(function = "validate_sampling_params"))]
11pub struct SamplingParams {
12    /// Temperature for sampling (must be >= 0.0, no upper limit)
13    #[validate(range(min = 0.0))]
14    pub temperature: Option<f32>,
15    /// Maximum number of new tokens to generate (must be >= 0)
16    #[validate(range(min = 0))]
17    pub max_new_tokens: Option<u32>,
18    /// Top-p nucleus sampling (0.0 < top_p <= 1.0)
19    #[validate(custom(function = "validate_top_p_value"))]
20    pub top_p: Option<f32>,
21    /// Top-k sampling (-1 to disable, or >= 1)
22    #[validate(custom(function = "validate_top_k_value"))]
23    pub top_k: Option<i32>,
24    #[validate(range(min = -2.0, max = 2.0))]
25    pub frequency_penalty: Option<f32>,
26    #[validate(range(min = -2.0, max = 2.0))]
27    pub presence_penalty: Option<f32>,
28    #[validate(range(min = 0.0, max = 2.0))]
29    pub repetition_penalty: Option<f32>,
30    pub stop: Option<StringOrArray>,
31    pub ignore_eos: Option<bool>,
32    pub skip_special_tokens: Option<bool>,
33    pub json_schema: Option<String>,
34    pub regex: Option<String>,
35    pub ebnf: Option<String>,
36    #[validate(range(min = 0.0, max = 1.0))]
37    pub min_p: Option<f32>,
38    /// Minimum number of new tokens (validated in schema function for cross-field check with max_new_tokens)
39    pub min_new_tokens: Option<u32>,
40    pub stop_token_ids: Option<Vec<u32>>,
41    pub no_stop_trim: Option<bool>,
42    pub n: Option<u32>,
43    pub sampling_seed: Option<u64>,
44    /// Custom parameters for engine-specific sampling behavior.
45    pub custom_params: Option<Map<String, Value>>,
46}
47
48// ============================================================================
49// Shared Validation Functions
50// ============================================================================
51
52/// Validates top_p: 0.0 < top_p <= 1.0 (can't use range validator for open interval)
53pub fn validate_top_p_value(top_p: f32) -> Result<(), validator::ValidationError> {
54    if !(top_p > 0.0 && top_p <= 1.0) {
55        return Err(validator::ValidationError::new(
56            "top_p must be in (0, 1] - greater than 0.0 and at most 1.0",
57        ));
58    }
59    Ok(())
60}
61
62/// Validates top_k: -1 (disabled) or >= 1 (special -1 case - can't use range validator)
63pub fn validate_top_k_value(top_k: i32) -> Result<(), validator::ValidationError> {
64    if top_k != -1 && top_k < 1 {
65        return Err(validator::ValidationError::new(
66            "top_k must be -1 (disabled) or at least 1",
67        ));
68    }
69    Ok(())
70}
71
72// ============================================================================
73// SamplingParams-Specific Validation
74// ============================================================================
75
76/// Validation function for SamplingParams - cross-field validation only
77fn validate_sampling_params(params: &SamplingParams) -> Result<(), validator::ValidationError> {
78    // 1. Cross-field validation: min_new_tokens <= max_new_tokens
79    if let (Some(min), Some(max)) = (params.min_new_tokens, params.max_new_tokens) {
80        if min > max {
81            return Err(validator::ValidationError::new(
82                "min_new_tokens cannot exceed max_new_tokens",
83            ));
84        }
85    }
86
87    // 2. Validate mutually exclusive structured output constraints
88    let constraint_count = [
89        params.regex.is_some(),
90        params.ebnf.is_some(),
91        params.json_schema.is_some(),
92    ]
93    .iter()
94    .filter(|&&x| x)
95    .count();
96
97    if constraint_count > 1 {
98        return Err(validator::ValidationError::new(
99            "only one of regex, ebnf, or json_schema can be set",
100        ));
101    }
102
103    Ok(())
104}