Skip to main content

rig_core/providers/ollama/
extension.rs

1//! Ollama's typed request options and reply extras. The shared section
2//! goes to both routes, the OpenAI-compatible `/v1` and the native
3//! `/api/chat`; the `"ollama.chat"` section, with the model parameters, goes
4//! to the native route only and is skipped on `/v1`.
5//!
6//! ```
7//! use rig_core::completion::CompletionRequest;
8//! use rig_core::providers::ollama::extension::{KeepAlive, OllamaOptions};
9//!
10//! let options = OllamaOptions::default()
11//!     .keep_alive(KeepAlive::duration("5m"))
12//!     .num_ctx(8192);
13//! let request = CompletionRequest::new("hi").provider_option(options);
14//! ```
15
16use serde::{Deserialize, Serialize};
17use serde_json::Value;
18
19use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
20use crate::message::Api;
21
22/// The `ollama` provider: its key, [`OllamaOptions`] and [`OllamaExtras`].
23#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
24pub struct OllamaExt;
25
26impl ProviderExtension for OllamaExt {
27    const PROVIDER: &'static str = super::PROVIDER_NAME;
28    type Options = OllamaOptions;
29    type Extras = OllamaExtras;
30}
31
32/// Ollama's request options.
33#[non_exhaustive]
34#[derive(Clone, Debug, Default, PartialEq, Serialize)]
35pub struct OllamaOptions {
36    /// The fields both routes read.
37    #[serde(rename = "*")]
38    pub shared: OllamaShared,
39    /// The fields only `/api/chat` reads.
40    #[serde(rename = "ollama.chat")]
41    pub chat: OllamaNative,
42}
43
44impl ExtensionOptions for OllamaOptions {
45    type Ext = OllamaExt;
46}
47
48impl OllamaOptions {
49    /// How long the daemon keeps the model loaded after the request
50    /// (`keep_alive`).
51    pub fn keep_alive(mut self, keep_alive: KeepAlive) -> Self {
52        self.shared.keep_alive = Some(keep_alive);
53        self
54    }
55
56    /// The context window, in tokens (`options.num_ctx`).
57    pub fn num_ctx(mut self, num_ctx: u32) -> Self {
58        self.chat.options.num_ctx = Some(num_ctx);
59        self
60    }
61
62    /// The prompt tokens kept when the context is shifted
63    /// (`options.num_keep`).
64    pub fn num_keep(mut self, num_keep: u32) -> Self {
65        self.chat.options.num_keep = Some(num_keep);
66        self
67    }
68
69    /// Sample from the `top_k` most likely tokens (`options.top_k`).
70    pub fn top_k(mut self, top_k: u32) -> Self {
71        self.chat.options.top_k = Some(top_k);
72        self
73    }
74
75    /// Drop tokens below `min_p` times the likeliest token's probability
76    /// (`options.min_p`).
77    pub fn min_p(mut self, min_p: f64) -> Self {
78        self.chat.options.min_p = Some(min_p);
79        self
80    }
81
82    /// Penalize repeated tokens (`options.repeat_penalty`).
83    pub fn repeat_penalty(mut self, penalty: f64) -> Self {
84        self.chat.options.repeat_penalty = Some(penalty);
85        self
86    }
87
88    /// How far back the repeat penalty looks, in tokens; `0` disables it
89    /// and `-1` is the context window (`options.repeat_last_n`).
90    pub fn repeat_last_n(mut self, last_n: i32) -> Self {
91        self.chat.options.repeat_last_n = Some(last_n);
92        self
93    }
94
95    /// The model layers offloaded to the GPU (`options.num_gpu`).
96    pub fn num_gpu(mut self, num_gpu: i32) -> Self {
97        self.chat.options.num_gpu = Some(num_gpu);
98        self
99    }
100
101    /// The CPU threads generation uses (`options.num_thread`).
102    pub fn num_thread(mut self, num_thread: u32) -> Self {
103        self.chat.options.num_thread = Some(num_thread);
104        self
105    }
106
107    /// Return each generated token's log probability (`logprobs`), read
108    /// back as [`OllamaExtras::logprobs`].
109    pub fn logprobs(mut self, logprobs: bool) -> Self {
110        self.chat.logprobs = Some(logprobs);
111        self
112    }
113
114    /// The most likely alternatives returned beside each token
115    /// (`top_logprobs`).
116    pub fn top_logprobs(mut self, top_logprobs: u32) -> Self {
117        self.chat.top_logprobs = Some(top_logprobs);
118        self
119    }
120}
121
122/// The fields both Ollama routes read.
123#[non_exhaustive]
124#[derive(Clone, Debug, Default, PartialEq, Serialize)]
125pub struct OllamaShared {
126    /// `keep_alive`.
127    #[serde(skip_serializing_if = "Option::is_none")]
128    pub keep_alive: Option<KeepAlive>,
129}
130
131/// How long the daemon keeps a model loaded after a request.
132#[non_exhaustive]
133#[derive(Clone, Debug, PartialEq, Serialize)]
134#[serde(untagged)]
135pub enum KeepAlive {
136    /// A duration string such as `"5m"`; a negative one keeps the model
137    /// loaded.
138    Duration(String),
139    /// A number of seconds; `0` unloads the model at once and a negative
140    /// number keeps it loaded.
141    Seconds(i64),
142}
143
144impl KeepAlive {
145    /// A duration string such as `"5m"` or `"1h"`.
146    pub fn duration(duration: impl Into<String>) -> Self {
147        Self::Duration(duration.into())
148    }
149
150    /// A number of seconds.
151    pub fn seconds(seconds: i64) -> Self {
152        Self::Seconds(seconds)
153    }
154}
155
156/// The fields only `/api/chat` reads.
157#[non_exhaustive]
158#[derive(Clone, Debug, Default, PartialEq, Serialize)]
159pub struct OllamaNative {
160    /// `options`: model parameters, beside the ones the request and its
161    /// generation options set.
162    #[serde(skip_serializing_if = "ModelOptions::is_empty")]
163    pub options: ModelOptions,
164    /// `logprobs`.
165    #[serde(skip_serializing_if = "Option::is_none")]
166    pub logprobs: Option<bool>,
167    /// `top_logprobs`.
168    #[serde(skip_serializing_if = "Option::is_none")]
169    pub top_logprobs: Option<u32>,
170}
171
172/// The model parameters `/api/chat` reads in `options`. `temperature`,
173/// `num_predict`, `top_p`, `seed` and `stop` are not here: the request and
174/// its generation options set them.
175#[non_exhaustive]
176#[derive(Clone, Debug, Default, PartialEq, Serialize)]
177pub struct ModelOptions {
178    /// `num_ctx`.
179    #[serde(skip_serializing_if = "Option::is_none")]
180    pub num_ctx: Option<u32>,
181    /// `num_keep`.
182    #[serde(skip_serializing_if = "Option::is_none")]
183    pub num_keep: Option<u32>,
184    /// `top_k`.
185    #[serde(skip_serializing_if = "Option::is_none")]
186    pub top_k: Option<u32>,
187    /// `min_p`.
188    #[serde(skip_serializing_if = "Option::is_none")]
189    pub min_p: Option<f64>,
190    /// `repeat_penalty`.
191    #[serde(skip_serializing_if = "Option::is_none")]
192    pub repeat_penalty: Option<f64>,
193    /// `repeat_last_n`.
194    #[serde(skip_serializing_if = "Option::is_none")]
195    pub repeat_last_n: Option<i32>,
196    /// `num_gpu`.
197    #[serde(skip_serializing_if = "Option::is_none")]
198    pub num_gpu: Option<i32>,
199    /// `num_thread`.
200    #[serde(skip_serializing_if = "Option::is_none")]
201    pub num_thread: Option<u32>,
202}
203
204impl ModelOptions {
205    /// Whether no parameter is set.
206    pub fn is_empty(&self) -> bool {
207        self == &Self::default()
208    }
209}
210
211/// The fields of an `/api/chat` reply Rig does not normalize, read from
212/// the whole reply or a stream's final record. A `/v1` reply carries only
213/// [`Self::model`] and leaves the rest `None`.
214#[non_exhaustive]
215#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
216#[serde(default)]
217pub struct OllamaExtras {
218    /// `model`, on both routes.
219    pub model: Option<String>,
220    /// `created_at`.
221    pub created_at: Option<String>,
222    /// `done_reason` as the daemon spells it (`stop`, `length`, ...).
223    pub done_reason: Option<String>,
224    /// `total_duration`, in nanoseconds.
225    pub total_duration: Option<u64>,
226    /// `load_duration`, in nanoseconds.
227    pub load_duration: Option<u64>,
228    /// `prompt_eval_duration`, in nanoseconds.
229    pub prompt_eval_duration: Option<u64>,
230    /// `eval_duration`, in nanoseconds.
231    pub eval_duration: Option<u64>,
232    /// `prompt_eval_count`.
233    pub prompt_eval_count: Option<u64>,
234    /// `prompt_eval_cached_count`: the prompt tokens read from the cache.
235    pub prompt_eval_cached_count: Option<u64>,
236    /// `eval_count`.
237    pub eval_count: Option<u64>,
238    /// `logprobs`, when [`OllamaOptions::logprobs`] asked for them.
239    pub logprobs: Option<Vec<Logprob>>,
240}
241
242/// One generated token's log probability.
243#[non_exhaustive]
244#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
245#[serde(default)]
246pub struct Logprob {
247    /// `token`.
248    pub token: Option<String>,
249    /// `logprob`.
250    pub logprob: Option<f64>,
251    /// `bytes`, the token's UTF-8 bytes.
252    pub bytes: Option<Vec<u8>>,
253    /// `top_logprobs`, the likeliest alternatives, when
254    /// [`OllamaOptions::top_logprobs`] asked for them.
255    pub top_logprobs: Option<Vec<Logprob>>,
256}
257
258impl ReplyExtras for OllamaExtras {
259    fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
260        Self::deserialize(raw)
261    }
262}
263
264#[cfg(test)]
265mod tests;