rig_core/providers/llamacpp/
extension.rs1use serde::{Deserialize, Serialize};
14use serde_json::{Map, Value};
15
16use crate::completion::provider_options::reply_field;
17use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
18use crate::message::Api;
19
20#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
22pub struct LlamaCppExt;
23
24impl ProviderExtension for LlamaCppExt {
25 const PROVIDER: &'static str = super::PROVIDER_NAME;
26 type Options = LlamaCppOptions;
27 type Extras = LlamaCppExtras;
28}
29
30#[non_exhaustive]
32#[derive(Clone, Debug, Default, PartialEq, Serialize)]
33pub struct LlamaCppOptions {
34 #[serde(rename = "*")]
36 pub shared: LlamaCppShared,
37}
38
39#[non_exhaustive]
41#[derive(Clone, Debug, Default, PartialEq, Serialize)]
42pub struct LlamaCppShared {
43 #[serde(skip_serializing_if = "Map::is_empty")]
45 pub chat_template_kwargs: Map<String, Value>,
46 #[serde(skip_serializing_if = "Option::is_none")]
48 pub reasoning_format: Option<ReasoningFormat>,
49 #[serde(skip_serializing_if = "Option::is_none")]
51 pub n_probs: Option<u32>,
52 #[serde(skip_serializing_if = "Vec::is_empty")]
54 pub samplers: Vec<String>,
55 #[serde(skip_serializing_if = "Option::is_none")]
57 pub top_k: Option<i32>,
58 #[serde(skip_serializing_if = "Option::is_none")]
60 pub min_p: Option<f64>,
61 #[serde(skip_serializing_if = "Option::is_none")]
63 pub typical_p: Option<f64>,
64 #[serde(skip_serializing_if = "Option::is_none")]
66 pub mirostat: Option<u8>,
67 #[serde(skip_serializing_if = "Option::is_none")]
69 pub mirostat_tau: Option<f64>,
70 #[serde(skip_serializing_if = "Option::is_none")]
72 pub mirostat_eta: Option<f64>,
73 #[serde(skip_serializing_if = "Option::is_none")]
75 pub id_slot: Option<i32>,
76 #[serde(skip_serializing_if = "Option::is_none")]
78 pub timings_per_token: Option<bool>,
79}
80
81#[non_exhaustive]
83#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
84#[serde(rename_all = "kebab-case")]
85pub enum ReasoningFormat {
86 None,
88 Deepseek,
90 DeepseekLegacy,
92 Auto,
94}
95
96impl LlamaCppOptions {
97 pub fn new() -> Self {
99 Self::default()
100 }
101
102 pub fn chat_template_kwarg(mut self, key: impl Into<String>, value: Value) -> Self {
104 self.shared.chat_template_kwargs.insert(key.into(), value);
105 self
106 }
107
108 pub fn reasoning_format(mut self, format: ReasoningFormat) -> Self {
110 self.shared.reasoning_format = Some(format);
111 self
112 }
113
114 pub fn n_probs(mut self, count: u32) -> Self {
116 self.shared.n_probs = Some(count);
117 self
118 }
119
120 pub fn samplers(mut self, samplers: impl IntoIterator<Item = impl Into<String>>) -> Self {
122 self.shared.samplers = samplers.into_iter().map(Into::into).collect();
123 self
124 }
125
126 pub fn top_k(mut self, k: i32) -> Self {
128 self.shared.top_k = Some(k);
129 self
130 }
131
132 pub fn min_p(mut self, p: f64) -> Self {
134 self.shared.min_p = Some(p);
135 self
136 }
137
138 pub fn typical_p(mut self, p: f64) -> Self {
140 self.shared.typical_p = Some(p);
141 self
142 }
143
144 pub fn mirostat(mut self, version: u8) -> Self {
146 self.shared.mirostat = Some(version);
147 self
148 }
149
150 pub fn mirostat_tau(mut self, tau: f64) -> Self {
152 self.shared.mirostat_tau = Some(tau);
153 self
154 }
155
156 pub fn mirostat_eta(mut self, eta: f64) -> Self {
158 self.shared.mirostat_eta = Some(eta);
159 self
160 }
161
162 pub fn id_slot(mut self, slot: i32) -> Self {
164 self.shared.id_slot = Some(slot);
165 self
166 }
167
168 pub fn timings_per_token(mut self, per_token: bool) -> Self {
170 self.shared.timings_per_token = Some(per_token);
171 self
172 }
173}
174
175impl ExtensionOptions for LlamaCppOptions {
176 type Ext = LlamaCppExt;
177}
178
179#[non_exhaustive]
181#[derive(Clone, Debug, Default, PartialEq)]
182pub struct LlamaCppExtras {
183 pub timings: Option<Timings>,
185}
186
187#[non_exhaustive]
189#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
190pub struct Timings {
191 #[serde(default)]
193 pub cache_n: Option<u64>,
194 #[serde(default)]
196 pub prompt_n: Option<u64>,
197 #[serde(default)]
199 pub prompt_ms: Option<f64>,
200 #[serde(default)]
202 pub prompt_per_token_ms: Option<f64>,
203 #[serde(default)]
205 pub prompt_per_second: Option<f64>,
206 #[serde(default)]
208 pub predicted_n: Option<u64>,
209 #[serde(default)]
211 pub predicted_ms: Option<f64>,
212 #[serde(default)]
214 pub predicted_per_token_ms: Option<f64>,
215 #[serde(default)]
217 pub predicted_per_second: Option<f64>,
218}
219
220impl ReplyExtras for LlamaCppExtras {
221 fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
222 Ok(Self {
223 timings: reply_field(raw, "/timings")?,
224 })
225 }
226}
227
228#[cfg(test)]
229mod tests;