rig_core/providers/ollama/
extension.rs1use serde::{Deserialize, Serialize};
17use serde_json::Value;
18
19use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
20use crate::message::Api;
21
22#[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#[non_exhaustive]
34#[derive(Clone, Debug, Default, PartialEq, Serialize)]
35pub struct OllamaOptions {
36 #[serde(rename = "*")]
38 pub shared: OllamaShared,
39 #[serde(rename = "ollama.chat")]
41 pub chat: OllamaNative,
42}
43
44impl ExtensionOptions for OllamaOptions {
45 type Ext = OllamaExt;
46}
47
48impl OllamaOptions {
49 pub fn keep_alive(mut self, keep_alive: KeepAlive) -> Self {
52 self.shared.keep_alive = Some(keep_alive);
53 self
54 }
55
56 pub fn num_ctx(mut self, num_ctx: u32) -> Self {
58 self.chat.options.num_ctx = Some(num_ctx);
59 self
60 }
61
62 pub fn num_keep(mut self, num_keep: u32) -> Self {
65 self.chat.options.num_keep = Some(num_keep);
66 self
67 }
68
69 pub fn top_k(mut self, top_k: u32) -> Self {
71 self.chat.options.top_k = Some(top_k);
72 self
73 }
74
75 pub fn min_p(mut self, min_p: f64) -> Self {
78 self.chat.options.min_p = Some(min_p);
79 self
80 }
81
82 pub fn repeat_penalty(mut self, penalty: f64) -> Self {
84 self.chat.options.repeat_penalty = Some(penalty);
85 self
86 }
87
88 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 pub fn num_gpu(mut self, num_gpu: i32) -> Self {
97 self.chat.options.num_gpu = Some(num_gpu);
98 self
99 }
100
101 pub fn num_thread(mut self, num_thread: u32) -> Self {
103 self.chat.options.num_thread = Some(num_thread);
104 self
105 }
106
107 pub fn logprobs(mut self, logprobs: bool) -> Self {
110 self.chat.logprobs = Some(logprobs);
111 self
112 }
113
114 pub fn top_logprobs(mut self, top_logprobs: u32) -> Self {
117 self.chat.top_logprobs = Some(top_logprobs);
118 self
119 }
120}
121
122#[non_exhaustive]
124#[derive(Clone, Debug, Default, PartialEq, Serialize)]
125pub struct OllamaShared {
126 #[serde(skip_serializing_if = "Option::is_none")]
128 pub keep_alive: Option<KeepAlive>,
129}
130
131#[non_exhaustive]
133#[derive(Clone, Debug, PartialEq, Serialize)]
134#[serde(untagged)]
135pub enum KeepAlive {
136 Duration(String),
139 Seconds(i64),
142}
143
144impl KeepAlive {
145 pub fn duration(duration: impl Into<String>) -> Self {
147 Self::Duration(duration.into())
148 }
149
150 pub fn seconds(seconds: i64) -> Self {
152 Self::Seconds(seconds)
153 }
154}
155
156#[non_exhaustive]
158#[derive(Clone, Debug, Default, PartialEq, Serialize)]
159pub struct OllamaNative {
160 #[serde(skip_serializing_if = "ModelOptions::is_empty")]
163 pub options: ModelOptions,
164 #[serde(skip_serializing_if = "Option::is_none")]
166 pub logprobs: Option<bool>,
167 #[serde(skip_serializing_if = "Option::is_none")]
169 pub top_logprobs: Option<u32>,
170}
171
172#[non_exhaustive]
176#[derive(Clone, Debug, Default, PartialEq, Serialize)]
177pub struct ModelOptions {
178 #[serde(skip_serializing_if = "Option::is_none")]
180 pub num_ctx: Option<u32>,
181 #[serde(skip_serializing_if = "Option::is_none")]
183 pub num_keep: Option<u32>,
184 #[serde(skip_serializing_if = "Option::is_none")]
186 pub top_k: Option<u32>,
187 #[serde(skip_serializing_if = "Option::is_none")]
189 pub min_p: Option<f64>,
190 #[serde(skip_serializing_if = "Option::is_none")]
192 pub repeat_penalty: Option<f64>,
193 #[serde(skip_serializing_if = "Option::is_none")]
195 pub repeat_last_n: Option<i32>,
196 #[serde(skip_serializing_if = "Option::is_none")]
198 pub num_gpu: Option<i32>,
199 #[serde(skip_serializing_if = "Option::is_none")]
201 pub num_thread: Option<u32>,
202}
203
204impl ModelOptions {
205 pub fn is_empty(&self) -> bool {
207 self == &Self::default()
208 }
209}
210
211#[non_exhaustive]
215#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
216#[serde(default)]
217pub struct OllamaExtras {
218 pub model: Option<String>,
220 pub created_at: Option<String>,
222 pub done_reason: Option<String>,
224 pub total_duration: Option<u64>,
226 pub load_duration: Option<u64>,
228 pub prompt_eval_duration: Option<u64>,
230 pub eval_duration: Option<u64>,
232 pub prompt_eval_count: Option<u64>,
234 pub prompt_eval_cached_count: Option<u64>,
236 pub eval_count: Option<u64>,
238 pub logprobs: Option<Vec<Logprob>>,
240}
241
242#[non_exhaustive]
244#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
245#[serde(default)]
246pub struct Logprob {
247 pub token: Option<String>,
249 pub logprob: Option<f64>,
251 pub bytes: Option<Vec<u8>>,
253 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;