use std::collections::HashMap;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use validator::Validate;
use super::{
common::{default_true, deserialize_null_as_false, is_false, GenerationRequest, InputIds},
sampling_params::SamplingParams,
};
use crate::validated::Normalizable;
#[derive(Clone, Debug, Serialize, Deserialize, Validate, schemars::JsonSchema)]
#[validate(schema(function = "validate_generate_request"))]
pub struct GenerateRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub text: Option<String>,
#[serde(default = "super::common::default_unknown_model")]
pub model: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub input_ids: Option<InputIds>,
#[serde(skip_serializing_if = "Option::is_none")]
pub input_embeds: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub image_data: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub video_data: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub audio_data: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub sampling_params: Option<SamplingParams>,
#[serde(skip_serializing_if = "Option::is_none")]
pub return_logprob: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logprob_start_len: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_logprobs_num: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub token_ids_logprob: Option<Vec<u32>>,
#[serde(default)]
pub return_text_in_logprobs: bool,
#[serde(default, deserialize_with = "deserialize_null_as_false")]
pub stream: bool,
#[serde(default = "default_true")]
pub log_metrics: bool,
#[serde(default, skip_serializing_if = "is_false")]
pub return_hidden_states: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub modalities: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub session_params: Option<HashMap<String, Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub lora_path: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub lora_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub custom_logit_processor: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub bootstrap_host: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub bootstrap_port: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub bootstrap_room: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub bootstrap_pair_key: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub data_parallel_rank: Option<i32>,
#[serde(default)]
pub background: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub conversation_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub priority: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub extra_key: Option<String>,
#[serde(default)]
pub no_logs: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub custom_labels: Option<HashMap<String, String>>,
#[serde(default)]
pub return_bytes: bool,
#[serde(default)]
pub return_entropy: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub rid: Option<String>,
#[serde(flatten)]
pub other: Map<String, Value>,
}
impl Normalizable for GenerateRequest {
}
fn validate_generate_request(req: &GenerateRequest) -> Result<(), validator::ValidationError> {
let has_text = req.text.is_some();
let has_input_ids = req.input_ids.is_some();
let count = [has_text, has_input_ids].iter().filter(|&&x| x).count();
if count == 0 {
return Err(validator::ValidationError::new(
"Either text or input_ids should be provided.",
));
}
if count > 1 {
return Err(validator::ValidationError::new(
"Either text or input_ids should be provided.",
));
}
Ok(())
}
impl GenerationRequest for GenerateRequest {
fn rid(&self) -> Option<&str> {
self.rid.as_deref()
}
fn is_stream(&self) -> bool {
self.stream
}
fn get_model(&self) -> Option<&str> {
Some(self.model.as_str())
}
fn extract_text_for_routing(&self) -> String {
if let Some(ref text) = self.text {
return text.clone();
}
if let Some(ref input_ids) = self.input_ids {
return match input_ids {
InputIds::Single(ids) => ids
.iter()
.map(|&id| id.to_string())
.collect::<Vec<String>>()
.join(" "),
InputIds::Batch(batches) => batches
.iter()
.flat_map(|batch| batch.iter().map(|&id| id.to_string()))
.collect::<Vec<String>>()
.join(" "),
};
}
String::new()
}
fn routing_tokens(&self) -> Option<&[i32]> {
match &self.input_ids {
Some(InputIds::Single(ids)) if !ids.is_empty() => Some(ids),
Some(InputIds::Batch(seqs)) => seqs
.first()
.map(Vec::as_slice)
.filter(|ids| !ids.is_empty()),
_ => None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
pub struct GenerateResponse {
pub text: String,
pub output_ids: Vec<u32>,
pub meta_info: GenerateMetaInfo,
}
#[serde_with::skip_serializing_none]
#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
pub struct GenerateMetaInfo {
pub id: String,
pub finish_reason: GenerateFinishReason,
pub prompt_tokens: u32,
pub weight_version: String,
pub input_token_logprobs: Option<Vec<Vec<Option<f64>>>>,
pub output_token_logprobs: Option<Vec<Vec<Option<f64>>>>,
pub completion_tokens: u32,
pub cached_tokens: u32,
pub reasoning_tokens: Option<u32>,
pub e2e_latency: f64,
pub matched_stop: Option<Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
#[serde(untagged)]
pub enum GenerateFinishReason {
Length {
#[serde(rename = "type")]
finish_type: GenerateFinishType,
length: u32,
},
Stop {
#[serde(rename = "type")]
finish_type: GenerateFinishType,
},
Other(Value),
}
#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
#[serde(rename_all = "lowercase")]
pub enum GenerateFinishType {
Length,
Stop,
}
#[cfg(test)]
mod tests {
use super::*;
fn req() -> GenerateRequest {
serde_json::from_value(serde_json::json!({"model": "m"})).expect("minimal request")
}
#[test]
fn return_hidden_states_false_is_omitted_and_absent_reads_false() {
let r = req();
let v = serde_json::to_value(&r).expect("serialize");
assert!(v.get("return_hidden_states").is_none());
let back: GenerateRequest = serde_json::from_value(v).expect("roundtrip");
assert!(!back.return_hidden_states);
}
#[test]
fn return_hidden_states_true_round_trips() {
let mut r = req();
r.return_hidden_states = true;
let v = serde_json::to_value(&r).expect("serialize");
assert_eq!(v["return_hidden_states"], true);
let back: GenerateRequest = serde_json::from_value(v).expect("roundtrip");
assert!(back.return_hidden_states);
}
#[test]
fn routing_tokens_from_single_input_ids() {
let mut r = req();
r.input_ids = Some(InputIds::Single(vec![1, 2, 3]));
assert_eq!(r.routing_tokens(), Some(&[1, 2, 3][..]));
}
#[test]
fn routing_tokens_prefer_input_ids_over_text() {
let mut r = req();
r.text = Some("hello".to_string());
r.input_ids = Some(InputIds::Single(vec![1, 2, 3]));
assert_eq!(r.routing_tokens(), Some(&[1, 2, 3][..]));
r.input_ids = Some(InputIds::Batch(vec![vec![4, 5], vec![6]]));
assert_eq!(r.routing_tokens(), Some(&[4, 5][..]));
}
#[test]
fn routing_tokens_none_for_empty_input_ids_with_text() {
let mut r = req();
r.text = Some("hello".to_string());
r.input_ids = Some(InputIds::Single(vec![]));
assert_eq!(r.routing_tokens(), None);
assert_eq!(r.extract_text_for_routing(), "hello");
}
#[test]
fn routing_tokens_from_batch_first_sequence() {
let mut r = req();
r.input_ids = Some(InputIds::Batch(vec![vec![1, 2], vec![3, 4]]));
assert_eq!(r.routing_tokens(), Some(&[1, 2][..]));
}
#[test]
fn routing_tokens_none_for_empty_inputs() {
let mut r = req();
r.input_ids = Some(InputIds::Batch(vec![]));
assert_eq!(r.routing_tokens(), None);
let mut r = req();
r.input_ids = Some(InputIds::Batch(vec![vec![], vec![1]]));
assert_eq!(r.routing_tokens(), None);
let r = req();
assert_eq!(r.routing_tokens(), None);
}
}