use serde::{Deserialize, Serialize};
use super::Dtype;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct Parameter {
pub name: String,
#[serde(default)]
pub dtype: Option<Dtype>,
#[serde(default)]
pub description: Option<String>,
#[serde(default)]
pub denotation: Option<String>,
#[serde(default)]
pub optional: Option<bool>,
#[serde(default)]
pub enumeration: Option<Vec<EnumerationMember>>,
#[serde(default)]
pub schema: Option<serde_json::Map<String, serde_json::Value>>,
#[serde(default)]
pub min: Option<f64>,
#[serde(default)]
pub max: Option<f64>,
#[serde(default)]
pub sample_rate: Option<u32>,
#[serde(default)]
pub context_length: Option<u32>,
#[serde(default)]
pub batch: Option<BatchConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum EnumerationValue {
String(String),
Int(i64),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EnumerationMember {
pub name: String,
pub value: EnumerationValue,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum BatchMode {
Static,
Dynamic,
Continuous,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct BatchConfig {
pub mode: BatchMode,
#[serde(default, alias = "max_count", alias = "maxCount", skip_serializing_if = "Option::is_none")]
pub capacity: Option<usize>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn context_length_round_trips_camel_case() {
let param: Parameter = serde_json::from_str(r#"{
"name": "messages",
"dtype": "list",
"denotation": "openai.chat.completions.messages",
"contextLength": 262144
}"#).unwrap();
assert_eq!(param.context_length, Some(262_144));
let json = serde_json::to_value(¶m).unwrap();
assert_eq!(json["contextLength"], 262_144);
}
#[test]
fn context_length_defaults_to_none() {
let param: Parameter = serde_json::from_str(r#"{ "name": "x", "dtype": "int32" }"#).unwrap();
assert!(param.context_length.is_none());
}
}