use std::fmt;
use std::net::SocketAddr;
use serde::{Deserialize, Serialize};
use crate::queue::QueueConfig;
mod imp;
mod interpolate;
mod validate;
#[cfg(test)]
pub(crate) use interpolate::interpolate;
pub(crate) use interpolate::interpolate_value;
#[cfg(test)]
use crate::error::ConfigError;
#[cfg(test)]
mod tests;
fn default_gpu_layers() -> u32 {
99
}
fn default_true() -> bool {
true
}
fn default_cache_type_k() -> String {
"q8_0".to_owned()
}
fn default_cache_type_v() -> String {
"q4_0".to_owned()
}
fn default_n_predict() -> u32 {
8192
}
#[derive(Clone)]
pub struct Secret(String);
impl Secret {
#[must_use]
pub(crate) fn new(value: String) -> Secret {
Secret(value)
}
#[must_use]
pub fn expose(&self) -> &str {
&self.0
}
#[must_use]
pub(crate) fn is_empty(&self) -> bool {
self.0.is_empty()
}
}
fn de_secret<'de, D>(deserializer: D) -> Result<Secret, D::Error>
where
D: serde::Deserializer<'de>,
{
let raw = String::deserialize(deserializer)?;
Ok(Secret::new(raw))
}
impl fmt::Debug for Secret {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("Secret(redacted)")
}
}
impl fmt::Display for Secret {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("redacted")
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
pub(crate) enum Protocol {
Openai,
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct Config {
pub(crate) server: ServerConfig,
pub(crate) queue: QueueConfig,
pub(crate) local: LocalConfig,
pub(crate) devices: Vec<DeviceConfig>,
pub(crate) endpoints: Vec<EndpointConfig>,
pub(crate) models: Vec<ModelConfig>,
pub(crate) local_models: Vec<LocalModelConfig>,
pub(crate) tools: Option<ToolsConfig>,
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
struct RawConfig {
server: ServerConfig,
#[serde(default)]
queue: QueueConfig,
#[serde(default)]
local: LocalConfig,
#[serde(rename = "device", default)]
devices: Vec<DeviceConfig>,
#[serde(rename = "endpoint", default)]
endpoints: Vec<EndpointConfig>,
#[serde(rename = "model", default)]
models: Vec<ModelConfig>,
#[serde(rename = "local_model", default)]
local_models: Vec<LocalModelConfig>,
#[serde(default)]
tools: Option<ToolsConfig>,
}
impl From<RawConfig> for Config {
fn from(raw: RawConfig) -> Config {
Config {
server: raw.server,
queue: raw.queue,
local: raw.local,
devices: raw.devices,
endpoints: raw.endpoints,
models: raw.models,
local_models: raw.local_models,
tools: raw.tools,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub(crate) enum DeviceKind {
Remote,
Local,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
#[non_exhaustive]
pub(crate) struct DeviceConfig {
pub id: String,
#[serde(rename = "type")]
pub kind: DeviceKind,
#[serde(default)]
pub concurrency: Option<usize>,
#[serde(default, rename = "lane")]
pub lanes: Vec<LaneConfig>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
#[non_exhaustive]
pub(crate) struct LaneConfig {
pub id: String,
pub concurrency: usize,
#[serde(default)]
pub device: Option<String>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
#[non_exhaustive]
pub(crate) struct LocalConfig {
#[serde(default)]
pub cache_dir: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
#[non_exhaustive]
pub(crate) struct LocalModelConfig {
pub name: String,
pub description: String,
pub source: String,
#[serde(default)]
pub sha256: Option<String>,
#[serde(default)]
pub device: Option<String>,
#[serde(default)]
pub lane: Option<String>,
pub context: u32,
#[serde(default)]
pub thinking: ThinkingMode,
#[serde(default = "default_gpu_layers")]
pub gpu_layers: u32,
#[serde(default = "default_true")]
pub flash_attention: bool,
#[serde(default = "default_cache_type_k")]
pub cache_type_k: String,
#[serde(default = "default_cache_type_v")]
pub cache_type_v: String,
#[serde(default = "default_n_predict")]
pub n_predict: u32,
#[serde(default)]
pub chat_template_file: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
#[non_exhaustive]
pub(crate) struct ServerConfig {
pub bind: SocketAddr,
#[serde(deserialize_with = "de_secret")]
pub key: Secret,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
#[non_exhaustive]
pub(crate) struct EndpointConfig {
pub id: String,
pub protocol: Protocol,
pub base_url: String,
#[serde(deserialize_with = "de_secret")]
pub api_key: Secret,
#[serde(default)]
pub concurrency: Option<usize>,
#[serde(default)]
pub device: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub(crate) enum ThinkingMode {
#[default]
Never,
Always,
Switchable,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
#[non_exhaustive]
pub(crate) struct ModelConfig {
pub name: String,
pub description: String,
pub context: u32,
#[serde(default)]
pub thinking: ThinkingMode,
pub upstream: String,
pub endpoints: Vec<String>,
#[serde(default)]
pub default_max_tokens: Option<u32>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
#[non_exhaustive]
pub(crate) struct ToolsConfig {
#[serde(default)]
pub web_search: Option<WebSearchConfig>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
#[non_exhaustive]
pub(crate) struct WebSearchConfig {
pub provider: SearchProvider,
#[serde(deserialize_with = "de_secret")]
pub api_key: Secret,
#[serde(default = "default_brave_base_url")]
pub base_url: String,
#[serde(default = "default_web_search_count")]
pub default_count: u8,
#[serde(default = "default_web_search_max_count")]
pub max_count: u8,
#[serde(default = "default_web_search_max_per_host")]
pub max_per_host: u8,
#[serde(default)]
pub default_freshness: String,
#[serde(default)]
pub default_safesearch: String,
#[serde(default = "default_true")]
pub strip_tracking: bool,
}
fn default_brave_base_url() -> String {
"https://api.search.brave.com/res/v1".to_string()
}
fn default_web_search_count() -> u8 {
10
}
fn default_web_search_max_count() -> u8 {
20
}
fn default_web_search_max_per_host() -> u8 {
2
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub(crate) enum SearchProvider {
Brave,
}
fn is_sha256_hex(value: &str) -> bool {
value.len() == 64 && value.chars().all(|c| matches!(c, '0'..='9' | 'a'..='f'))
}