use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use crate::completion::provider_options::reply_field;
use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
use crate::message::Api;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct LlamaCppExt;
impl ProviderExtension for LlamaCppExt {
const PROVIDER: &'static str = super::PROVIDER_NAME;
type Options = LlamaCppOptions;
type Extras = LlamaCppExtras;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct LlamaCppOptions {
#[serde(rename = "*")]
pub shared: LlamaCppShared,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct LlamaCppShared {
#[serde(skip_serializing_if = "Map::is_empty")]
pub chat_template_kwargs: Map<String, Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_format: Option<ReasoningFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub n_probs: Option<u32>,
#[serde(skip_serializing_if = "Vec::is_empty")]
pub samplers: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_k: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub min_p: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub typical_p: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub mirostat: Option<u8>,
#[serde(skip_serializing_if = "Option::is_none")]
pub mirostat_tau: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub mirostat_eta: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub id_slot: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub timings_per_token: Option<bool>,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "kebab-case")]
pub enum ReasoningFormat {
None,
Deepseek,
DeepseekLegacy,
Auto,
}
impl LlamaCppOptions {
pub fn new() -> Self {
Self::default()
}
pub fn chat_template_kwarg(mut self, key: impl Into<String>, value: Value) -> Self {
self.shared.chat_template_kwargs.insert(key.into(), value);
self
}
pub fn reasoning_format(mut self, format: ReasoningFormat) -> Self {
self.shared.reasoning_format = Some(format);
self
}
pub fn n_probs(mut self, count: u32) -> Self {
self.shared.n_probs = Some(count);
self
}
pub fn samplers(mut self, samplers: impl IntoIterator<Item = impl Into<String>>) -> Self {
self.shared.samplers = samplers.into_iter().map(Into::into).collect();
self
}
pub fn top_k(mut self, k: i32) -> Self {
self.shared.top_k = Some(k);
self
}
pub fn min_p(mut self, p: f64) -> Self {
self.shared.min_p = Some(p);
self
}
pub fn typical_p(mut self, p: f64) -> Self {
self.shared.typical_p = Some(p);
self
}
pub fn mirostat(mut self, version: u8) -> Self {
self.shared.mirostat = Some(version);
self
}
pub fn mirostat_tau(mut self, tau: f64) -> Self {
self.shared.mirostat_tau = Some(tau);
self
}
pub fn mirostat_eta(mut self, eta: f64) -> Self {
self.shared.mirostat_eta = Some(eta);
self
}
pub fn id_slot(mut self, slot: i32) -> Self {
self.shared.id_slot = Some(slot);
self
}
pub fn timings_per_token(mut self, per_token: bool) -> Self {
self.shared.timings_per_token = Some(per_token);
self
}
}
impl ExtensionOptions for LlamaCppOptions {
type Ext = LlamaCppExt;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq)]
pub struct LlamaCppExtras {
pub timings: Option<Timings>,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
pub struct Timings {
#[serde(default)]
pub cache_n: Option<u64>,
#[serde(default)]
pub prompt_n: Option<u64>,
#[serde(default)]
pub prompt_ms: Option<f64>,
#[serde(default)]
pub prompt_per_token_ms: Option<f64>,
#[serde(default)]
pub prompt_per_second: Option<f64>,
#[serde(default)]
pub predicted_n: Option<u64>,
#[serde(default)]
pub predicted_ms: Option<f64>,
#[serde(default)]
pub predicted_per_token_ms: Option<f64>,
#[serde(default)]
pub predicted_per_second: Option<f64>,
}
impl ReplyExtras for LlamaCppExtras {
fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
Ok(Self {
timings: reply_field(raw, "/timings")?,
})
}
}
#[cfg(test)]
mod tests;