use serde::{Deserialize, Serialize};
fn default_chat_template() -> String {
"chatml".into()
}
fn default_temperature() -> f64 {
0.7
}
fn default_max_tokens() -> usize {
2048
}
fn default_seed() -> u64 {
42
}
fn default_repeat_penalty() -> f32 {
1.1
}
fn default_repeat_last_n() -> usize {
64
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum CandleSource {
#[default]
Huggingface,
Local,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum CandleDevice {
#[default]
Cpu,
Cuda,
Metal,
Auto,
}
#[derive(Deserialize, Serialize)]
pub struct CandleConfig {
#[serde(default)]
pub source: CandleSource,
#[serde(default)]
pub local_path: String,
#[serde(default)]
pub filename: Option<String>,
#[serde(default = "default_chat_template")]
pub chat_template: String,
#[serde(default)]
pub device: CandleDevice,
#[serde(default)]
pub embedding_repo: Option<String>,
#[serde(default)]
pub hf_token: Option<String>,
#[serde(default)]
pub generation: GenerationParams,
#[serde(default = "default_inference_timeout_secs")]
pub inference_timeout_secs: u64,
}
fn default_inference_timeout_secs() -> u64 {
120
}
impl std::fmt::Debug for CandleConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CandleConfig")
.field("source", &self.source)
.field("local_path", &self.local_path)
.field("filename", &self.filename)
.field("chat_template", &self.chat_template)
.field("device", &self.device)
.field("embedding_repo", &self.embedding_repo)
.field("hf_token", &self.hf_token.as_ref().map(|_| "[REDACTED]"))
.field("generation", &self.generation)
.field("inference_timeout_secs", &self.inference_timeout_secs)
.finish()
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct GenerationParams {
#[serde(default = "default_temperature")]
pub temperature: f64,
#[serde(default)]
pub top_p: Option<f64>,
#[serde(default)]
pub top_k: Option<usize>,
#[serde(default = "default_max_tokens")]
pub max_tokens: usize,
#[serde(default = "default_seed")]
pub seed: u64,
#[serde(default = "default_repeat_penalty")]
pub repeat_penalty: f32,
#[serde(default = "default_repeat_last_n")]
pub repeat_last_n: usize,
}
pub const MAX_TOKENS_CAP: usize = 32768;
impl GenerationParams {
#[must_use]
pub fn capped_max_tokens(&self) -> usize {
self.max_tokens.min(MAX_TOKENS_CAP)
}
}
impl Default for GenerationParams {
fn default() -> Self {
Self {
temperature: default_temperature(),
top_p: None,
top_k: None,
max_tokens: default_max_tokens(),
seed: default_seed(),
repeat_penalty: default_repeat_penalty(),
repeat_last_n: default_repeat_last_n(),
}
}
}
#[derive(Clone, Deserialize, Serialize)]
pub struct CandleInlineConfig {
#[serde(default)]
pub source: CandleSource,
#[serde(default)]
pub local_path: String,
#[serde(default)]
pub filename: Option<String>,
#[serde(default)]
pub chat_model_sha256: Option<String>,
#[serde(default = "default_chat_template")]
pub chat_template: String,
#[serde(default)]
pub device: CandleDevice,
#[serde(default)]
pub embedding_repo: Option<String>,
#[serde(default)]
pub embedding_model_sha256: Option<String>,
#[serde(default)]
pub hf_token: Option<String>,
#[serde(default)]
pub generation: GenerationParams,
#[serde(default = "default_inference_timeout_secs")]
pub inference_timeout_secs: u64,
}
impl Default for CandleInlineConfig {
fn default() -> Self {
Self {
source: CandleSource::default(),
local_path: String::new(),
filename: None,
chat_model_sha256: None,
chat_template: default_chat_template(),
device: CandleDevice::default(),
embedding_repo: None,
embedding_model_sha256: None,
hf_token: None,
generation: GenerationParams::default(),
inference_timeout_secs: default_inference_timeout_secs(),
}
}
}
impl std::fmt::Debug for CandleInlineConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CandleInlineConfig")
.field("source", &self.source)
.field("local_path", &self.local_path)
.field("filename", &self.filename)
.field("chat_model_sha256", &self.chat_model_sha256)
.field("chat_template", &self.chat_template)
.field("device", &self.device)
.field("embedding_repo", &self.embedding_repo)
.field("embedding_model_sha256", &self.embedding_model_sha256)
.field("hf_token", &self.hf_token.as_ref().map(|_| "[REDACTED]"))
.field("generation", &self.generation)
.field("inference_timeout_secs", &self.inference_timeout_secs)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn candle_config_debug_redacts_hf_token() {
let cfg = CandleConfig {
source: CandleSource::default(),
local_path: String::new(),
filename: None,
chat_template: default_chat_template(),
device: CandleDevice::default(),
embedding_repo: None,
hf_token: Some("hf_SUPERSECRET".to_owned()),
generation: GenerationParams::default(),
inference_timeout_secs: default_inference_timeout_secs(),
};
let dbg = format!("{cfg:?}");
assert!(!dbg.contains("hf_SUPERSECRET"));
assert!(dbg.contains("[REDACTED]"));
}
#[test]
fn candle_config_debug_none_hf_token() {
let cfg = CandleConfig {
source: CandleSource::default(),
local_path: String::new(),
filename: None,
chat_template: default_chat_template(),
device: CandleDevice::default(),
embedding_repo: None,
hf_token: None,
generation: GenerationParams::default(),
inference_timeout_secs: default_inference_timeout_secs(),
};
let dbg = format!("{cfg:?}");
assert!(!dbg.contains("[REDACTED]"));
assert!(dbg.contains("hf_token: None"));
}
#[test]
fn candle_inline_config_debug_redacts_hf_token() {
let cfg = CandleInlineConfig {
hf_token: Some("hf_SUPERSECRET".to_owned()),
..CandleInlineConfig::default()
};
let dbg = format!("{cfg:?}");
assert!(!dbg.contains("hf_SUPERSECRET"));
assert!(dbg.contains("[REDACTED]"));
}
#[test]
fn candle_inline_config_debug_none_hf_token() {
let cfg = CandleInlineConfig::default();
let dbg = format!("{cfg:?}");
assert!(!dbg.contains("[REDACTED]"));
assert!(dbg.contains("hf_token: None"));
}
}