Skip to main content

rig_core/providers/llamacpp/
extension.rs

1//! llama.cpp's typed request options and reply extras
2//! (<https://github.com/ggml-org/llama.cpp/blob/master/tools/server/README.md>).
3//!
4//! ```
5//! use rig_core::completion::CompletionRequest;
6//! use rig_core::providers::llamacpp::extension::{LlamaCppOptions};
7//!
8//! let options = LlamaCppOptions::new().top_k(40).min_p(0.05);
9//! let request = CompletionRequest::new("hi").provider_option(options);
10//! # let _ = request;
11//! ```
12
13use 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/// llama.cpp's extension marker.
21#[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/// llama.cpp's request options.
31#[non_exhaustive]
32#[derive(Clone, Debug, Default, PartialEq, Serialize)]
33pub struct LlamaCppOptions {
34    /// The fields every route takes.
35    #[serde(rename = "*")]
36    pub shared: LlamaCppShared,
37}
38
39/// The fields `llama-server` takes.
40#[non_exhaustive]
41#[derive(Clone, Debug, Default, PartialEq, Serialize)]
42pub struct LlamaCppShared {
43    /// Arguments to the model's chat template, such as `enable_thinking`.
44    #[serde(skip_serializing_if = "Map::is_empty")]
45    pub chat_template_kwargs: Map<String, Value>,
46    /// How the reply carries the reasoning.
47    #[serde(skip_serializing_if = "Option::is_none")]
48    pub reasoning_format: Option<ReasoningFormat>,
49    /// How many of the most likely tokens to report per position.
50    #[serde(skip_serializing_if = "Option::is_none")]
51    pub n_probs: Option<u32>,
52    /// The samplers, in the order they run.
53    #[serde(skip_serializing_if = "Vec::is_empty")]
54    pub samplers: Vec<String>,
55    /// Sample from the `k` most likely tokens.
56    #[serde(skip_serializing_if = "Option::is_none")]
57    pub top_k: Option<i32>,
58    /// The minimum probability of a token, relative to the most likely one.
59    #[serde(skip_serializing_if = "Option::is_none")]
60    pub min_p: Option<f64>,
61    /// Locally typical sampling.
62    #[serde(skip_serializing_if = "Option::is_none")]
63    pub typical_p: Option<f64>,
64    /// Mirostat sampling: 0 off, 1 Mirostat, 2 Mirostat 2.0.
65    #[serde(skip_serializing_if = "Option::is_none")]
66    pub mirostat: Option<u8>,
67    /// Mirostat's target entropy.
68    #[serde(skip_serializing_if = "Option::is_none")]
69    pub mirostat_tau: Option<f64>,
70    /// Mirostat's learning rate.
71    #[serde(skip_serializing_if = "Option::is_none")]
72    pub mirostat_eta: Option<f64>,
73    /// The server slot that runs the request.
74    #[serde(skip_serializing_if = "Option::is_none")]
75    pub id_slot: Option<i32>,
76    /// Whether every streamed chunk carries timings.
77    #[serde(skip_serializing_if = "Option::is_none")]
78    pub timings_per_token: Option<bool>,
79}
80
81/// How a reply carries the reasoning.
82#[non_exhaustive]
83#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
84#[serde(rename_all = "kebab-case")]
85pub enum ReasoningFormat {
86    /// Left in the content.
87    None,
88    /// In `reasoning_content`.
89    Deepseek,
90    /// In `reasoning_content`, and in the content's `<think>` tags too.
91    DeepseekLegacy,
92    /// As the server's template decides.
93    Auto,
94}
95
96impl LlamaCppOptions {
97    /// No option set.
98    pub fn new() -> Self {
99        Self::default()
100    }
101
102    /// Pass `key`, `value` to the model's chat template.
103    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    /// Carry the reasoning as `format`.
109    pub fn reasoning_format(mut self, format: ReasoningFormat) -> Self {
110        self.shared.reasoning_format = Some(format);
111        self
112    }
113
114    /// Report the `count` most likely tokens per position.
115    pub fn n_probs(mut self, count: u32) -> Self {
116        self.shared.n_probs = Some(count);
117        self
118    }
119
120    /// Run `samplers`, in order.
121    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    /// Sample from the `k` most likely tokens.
127    pub fn top_k(mut self, k: i32) -> Self {
128        self.shared.top_k = Some(k);
129        self
130    }
131
132    /// Set the minimum relative token probability.
133    pub fn min_p(mut self, p: f64) -> Self {
134        self.shared.min_p = Some(p);
135        self
136    }
137
138    /// Set locally typical sampling.
139    pub fn typical_p(mut self, p: f64) -> Self {
140        self.shared.typical_p = Some(p);
141        self
142    }
143
144    /// Set the Mirostat version: 0 off, 1 or 2.
145    pub fn mirostat(mut self, version: u8) -> Self {
146        self.shared.mirostat = Some(version);
147        self
148    }
149
150    /// Set Mirostat's target entropy.
151    pub fn mirostat_tau(mut self, tau: f64) -> Self {
152        self.shared.mirostat_tau = Some(tau);
153        self
154    }
155
156    /// Set Mirostat's learning rate.
157    pub fn mirostat_eta(mut self, eta: f64) -> Self {
158        self.shared.mirostat_eta = Some(eta);
159        self
160    }
161
162    /// Run the request on server slot `slot`.
163    pub fn id_slot(mut self, slot: i32) -> Self {
164        self.shared.id_slot = Some(slot);
165        self
166    }
167
168    /// Whether every streamed chunk carries timings.
169    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/// llama.cpp's reply fields. Each is `None` when the reply lacks it.
180#[non_exhaustive]
181#[derive(Clone, Debug, Default, PartialEq)]
182pub struct LlamaCppExtras {
183    /// How long the prompt and the completion took.
184    pub timings: Option<Timings>,
185}
186
187/// How long a `llama-server` request took.
188#[non_exhaustive]
189#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
190pub struct Timings {
191    /// Prompt tokens reused from the cache.
192    #[serde(default)]
193    pub cache_n: Option<u64>,
194    /// Prompt tokens evaluated.
195    #[serde(default)]
196    pub prompt_n: Option<u64>,
197    /// Milliseconds spent on the prompt.
198    #[serde(default)]
199    pub prompt_ms: Option<f64>,
200    /// Milliseconds per prompt token.
201    #[serde(default)]
202    pub prompt_per_token_ms: Option<f64>,
203    /// Prompt tokens per second.
204    #[serde(default)]
205    pub prompt_per_second: Option<f64>,
206    /// Tokens generated.
207    #[serde(default)]
208    pub predicted_n: Option<u64>,
209    /// Milliseconds spent generating.
210    #[serde(default)]
211    pub predicted_ms: Option<f64>,
212    /// Milliseconds per generated token.
213    #[serde(default)]
214    pub predicted_per_token_ms: Option<f64>,
215    /// Generated tokens per second.
216    #[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;