use crate::registry::ModelInfo;
use serde::{Deserialize, Serialize};
#[derive(Serialize, Deserialize)]
pub struct HealthResponse {
pub status: String,
pub version: String,
pub compute_mode: String,
pub model_loaded: bool,
pub uptime_sec: f64,
}
#[derive(Serialize, Deserialize)]
pub struct TokenizeRequest {
pub text: String,
pub model_id: Option<String>,
}
#[derive(Serialize, Deserialize)]
pub struct TokenizeResponse {
pub token_ids: Vec<u32>,
pub num_tokens: usize,
}
#[derive(Serialize, Deserialize)]
pub struct GenerateRequest {
pub prompt: String,
#[serde(default = "default_max_tokens")]
pub max_tokens: usize,
#[serde(
default = "default_temperature",
deserialize_with = "deserialize_temperature_f32_required"
)]
pub temperature: f32,
#[serde(default = "default_strategy")]
pub strategy: String,
#[serde(default = "default_top_k")]
pub top_k: usize,
#[serde(default = "default_top_p")]
pub top_p: f32,
pub seed: Option<u64>,
pub model_id: Option<String>,
}
pub fn default_max_tokens() -> usize {
50
}
pub(crate) fn default_temperature() -> f32 {
1.0
}
pub(crate) fn default_strategy() -> String {
"greedy".to_string()
}
pub fn default_top_k() -> usize {
50
}
pub(crate) fn default_top_p() -> f32 {
0.9
}
#[derive(Serialize, Deserialize)]
pub struct GenerateResponse {
pub token_ids: Vec<u32>,
pub text: String,
pub num_generated: usize,
}
#[derive(Serialize, Deserialize)]
pub struct ErrorResponse {
pub error: String,
}
#[derive(Serialize, Deserialize)]
pub struct BatchTokenizeRequest {
pub texts: Vec<String>,
}
#[derive(Serialize, Deserialize)]
pub struct BatchTokenizeResponse {
pub results: Vec<TokenizeResponse>,
}
#[derive(Serialize, Deserialize)]
pub struct BatchGenerateRequest {
pub prompts: Vec<String>,
#[serde(default = "default_max_tokens")]
pub max_tokens: usize,
#[serde(
default = "default_temperature",
deserialize_with = "deserialize_temperature_f32_required"
)]
pub temperature: f32,
#[serde(default = "default_strategy")]
pub strategy: String,
#[serde(default = "default_top_k")]
pub top_k: usize,
#[serde(default = "default_top_p")]
pub top_p: f32,
pub seed: Option<u64>,
}
#[derive(Serialize, Deserialize)]
pub struct BatchGenerateResponse {
pub results: Vec<GenerateResponse>,
}
#[derive(Serialize, Deserialize)]
pub struct StreamTokenEvent {
pub token_id: u32,
pub text: String,
}
#[derive(Serialize, Deserialize)]
pub struct StreamDoneEvent {
pub num_generated: usize,
}
#[derive(Serialize, Deserialize)]
pub struct ModelsResponse {
pub models: Vec<ModelInfo>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FinishReason {
Stop,
Length,
}
impl FinishReason {
#[must_use]
pub fn from_generation(stopped: bool, completion_tokens: usize, max_tokens: usize) -> Self {
if !stopped && completion_tokens >= max_tokens {
Self::Length
} else {
Self::Stop
}
}
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
Self::Stop => "stop",
Self::Length => "length",
}
}
}
impl std::fmt::Display for FinishReason {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(transparent)]
pub struct ChoiceCount(usize);
impl ChoiceCount {
pub const ONE: Self = Self(1);
#[must_use]
pub fn get(self) -> usize {
self.0
}
}
impl Default for ChoiceCount {
fn default() -> Self {
Self::ONE
}
}
impl<'de> Deserialize<'de> for ChoiceCount {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let requested = usize::deserialize(deserializer)?;
if requested == 1 {
Ok(Self::ONE)
} else {
Err(serde::de::Error::custom(format!(
"{}n must be 1: this server returns exactly one choice per request, \
so n={requested} cannot be honoured (send {requested} requests instead)",
crate::api::CLIENT_VISIBLE_MARKER
)))
}
}
}
#[must_use]
pub(crate) fn temperature_is_servable(temperature: f64) -> bool {
temperature.is_finite() && temperature >= 0.0
}
fn temperature_rejection<E: serde::de::Error>(temperature: f64) -> E {
E::custom(format!(
"{}temperature must be a finite number >= 0 (0 means deterministic/greedy), \
got {temperature}",
crate::api::CLIENT_VISIBLE_MARKER
))
}
pub(crate) fn deserialize_temperature_f32<'de, D>(deserializer: D) -> Result<Option<f32>, D::Error>
where
D: serde::Deserializer<'de>,
{
let Some(raw) = Option::<f64>::deserialize(deserializer)? else {
return Ok(None);
};
let narrowed = f64::from(raw as f32);
if !temperature_is_servable(raw) || !temperature_is_servable(narrowed) {
return Err(temperature_rejection(raw));
}
Ok(Some(raw as f32))
}
pub(crate) fn deserialize_temperature_f64<'de, D>(deserializer: D) -> Result<Option<f64>, D::Error>
where
D: serde::Deserializer<'de>,
{
Ok(deserialize_temperature_f32(deserializer)?.map(f64::from))
}
pub(crate) fn deserialize_temperature_f32_required<'de, D>(deserializer: D) -> Result<f32, D::Error>
where
D: serde::Deserializer<'de>,
{
deserialize_temperature_f32(deserializer)?.ok_or_else(|| {
serde::de::Error::custom(format!(
"{}temperature must be a finite number >= 0, not null",
crate::api::CLIENT_VISIBLE_MARKER
))
})
}