1use crate::{FirstTokenDistribution, FirstTokenRequest, HostCapabilities, HostError, ModelHost};
4use async_trait::async_trait;
5use serde::Serialize;
6use serde_json::Value;
7use std::collections::HashSet;
8use std::sync::Mutex;
9use std::time::{Duration, Instant};
10
11const OLLAMA_MAX_TOP_LOGPROBS: usize = 20;
14
15#[derive(Debug, Clone, Copy, Serialize)]
29struct SamplingOptions {
30 temperature: f64,
31 top_k: u32,
32 top_p: f64,
33 min_p: f64,
34 num_predict: u32,
35}
36
37impl SamplingOptions {
38 const REQUIRED: Self = Self {
39 temperature: 1.0,
40 top_k: 0,
41 top_p: 1.0,
42 min_p: 0.0,
43 num_predict: 1,
44 };
45}
46
47#[derive(Debug, Serialize)]
48struct ChatRequest<'a> {
49 model: &'a str,
50 messages: Vec<Message<'a>>,
51 stream: bool,
52 #[serde(skip_serializing_if = "Option::is_none")]
55 think: Option<bool>,
56 logprobs: bool,
57 top_logprobs: usize,
58 options: SamplingOptions,
59 #[serde(skip_serializing_if = "Option::is_none")]
60 keep_alive: Option<&'a str>,
61}
62
63#[derive(Debug, Serialize)]
64struct Message<'a> {
65 role: &'a str,
66 content: &'a str,
67}
68
69pub struct OllamaHost {
70 client: reqwest::Client,
71 base_url: String,
72 timeout: Duration,
73 refused_thinking: Mutex<HashSet<String>>,
76}
77
78impl OllamaHost {
79 pub fn new(base_url: impl Into<String>, timeout: Duration) -> Result<Self, HostError> {
80 let client = reqwest::Client::builder()
81 .timeout(timeout)
82 .build()
83 .map_err(|e| HostError::Unavailable(e.to_string()))?;
84 Ok(Self {
85 client,
86 base_url: crate::base_url(&base_url.into())?,
87 timeout,
88 refused_thinking: Mutex::default(),
89 })
90 }
91
92 fn body<'a>(&self, req: &'a FirstTokenRequest, think: Option<bool>) -> ChatRequest<'a> {
93 let mut messages = Vec::with_capacity(2);
94 if let Some(system) = req.system.as_deref() {
95 messages.push(Message {
96 role: "system",
97 content: system,
98 });
99 }
100 messages.push(Message {
101 role: "user",
102 content: &req.user,
103 });
104
105 ChatRequest {
106 model: &req.model,
107 messages,
108 stream: false,
109 think,
110 logprobs: true,
111 top_logprobs: req.top_logprobs.min(OLLAMA_MAX_TOP_LOGPROBS),
112 options: SamplingOptions::REQUIRED,
113 keep_alive: req.keep_alive.as_deref(),
114 }
115 }
116
117 async fn post(&self, body: &ChatRequest<'_>) -> Result<(u16, String), HostError> {
118 let response = self
119 .client
120 .post(format!("{}/api/chat", self.base_url))
121 .json(body)
122 .send()
123 .await
124 .map_err(|e| HostError::transport(&e, self.timeout))?;
125
126 let status = response.status().as_u16();
127 let text = response
128 .text()
129 .await
130 .map_err(|e| HostError::transport(&e, self.timeout))?;
131 Ok((status, text))
132 }
133}
134
135fn rejects_thinking(body: &str) -> bool {
143 let lower = body.to_ascii_lowercase();
144 lower.contains("not support") && lower.contains("thinking")
145}
146
147#[async_trait]
148impl ModelHost for OllamaHost {
149 fn capabilities(&self) -> HostCapabilities {
150 HostCapabilities {
151 max_top_logprobs: OLLAMA_MAX_TOP_LOGPROBS,
152 }
153 }
154
155 fn name(&self) -> &str {
156 "ollama"
157 }
158
159 async fn first_token(
160 &self,
161 req: FirstTokenRequest,
162 ) -> Result<FirstTokenDistribution, HostError> {
163 let started = Instant::now();
164
165 let known_to_reject = self
166 .refused_thinking
167 .lock()
168 .expect("lock is never poisoned")
169 .contains(&req.model);
170 let think = if known_to_reject { None } else { Some(false) };
171
172 let (mut status, mut text) = self.post(&self.body(&req, think)).await?;
173
174 if think.is_some() && status >= 400 && rejects_thinking(&text) {
175 tracing::debug!(
176 model = %req.model,
177 "model rejects the think flag; retrying without it, and from now on"
178 );
179 self.refused_thinking
180 .lock()
181 .expect("lock is never poisoned")
182 .insert(req.model.clone());
183 (status, text) = self.post(&self.body(&req, None)).await?;
184 }
185
186 if status >= 400 {
187 return Err(HostError::status(status, text));
188 }
189
190 parse_distribution(&text, &req.model, started.elapsed())
191 }
192}
193
194fn parse_distribution(
196 text: &str,
197 model: &str,
198 latency: Duration,
199) -> Result<FirstTokenDistribution, HostError> {
200 let value: Value =
201 serde_json::from_str(text).map_err(|e| HostError::Protocol(e.to_string()))?;
202
203 let first = value
204 .get("logprobs")
205 .and_then(Value::as_array)
206 .and_then(|positions| positions.first())
207 .ok_or_else(|| HostError::NoLogprobs {
208 model: model.to_string(),
209 })?;
210
211 let ranked = first
213 .get("top_logprobs")
214 .and_then(Value::as_array)
215 .ok_or_else(|| HostError::NoLogprobs {
216 model: model.to_string(),
217 })?;
218
219 let mut tokens: Vec<(String, f64)> = ranked
220 .iter()
221 .filter_map(|entry| {
222 let token = entry.get("token")?.as_str()?.to_string();
223 let logprob = entry.get("logprob")?.as_f64()?;
224 Some((token, logprob))
225 })
226 .collect();
227
228 if tokens.is_empty() {
229 return Err(HostError::NoLogprobs {
230 model: model.to_string(),
231 });
232 }
233
234 tokens.sort_by(|a, b| b.1.total_cmp(&a.1));
235
236 let floor = tokens
237 .last()
238 .map(|(_, lp)| *lp)
239 .expect("tokens is non-empty");
240
241 let input_tokens = value
242 .get("prompt_eval_count")
243 .and_then(Value::as_u64)
244 .unwrap_or(0) as u32;
245
246 Ok(FirstTokenDistribution {
247 tokens,
248 floor,
249 input_tokens,
250 latency,
251 })
252}
253
254#[cfg(test)]
255mod tests {
256 use super::*;
257
258 const REAL_RESPONSE: &str = r#"{
260 "model": "gemma4:26b-a4b-it-q4_K_M",
261 "message": {"role": "assistant", "content": "D"},
262 "done": true,
263 "prompt_eval_count": 110,
264 "logprobs": [{
265 "token": "D",
266 "logprob": -0.005,
267 "top_logprobs": [
268 {"token": "D", "logprob": -0.005},
269 {"token": "A", "logprob": -5.246},
270 {"token": "C", "logprob": -11.515},
271 {"token": "B", "logprob": -13.662}
272 ]
273 }]
274 }"#;
275
276 #[test]
277 fn parses_ranked_tokens_and_floor() {
278 let dist = parse_distribution(REAL_RESPONSE, "m", Duration::from_millis(7)).unwrap();
279
280 assert_eq!(dist.tokens.len(), 4);
281 assert_eq!(dist.tokens[0].0, "D");
282 assert_eq!(dist.input_tokens, 110);
283 assert_eq!(dist.logprob("A"), Some(-5.246));
284 assert_eq!(dist.logprob("Z"), None);
285 assert_eq!(dist.floor, -13.662);
287 }
288
289 #[test]
290 fn ranks_tokens_even_when_the_host_reports_them_unsorted() {
291 let shuffled = r#"{"logprobs":[{"token":"B","top_logprobs":[
292 {"token":"A","logprob":-5.0},{"token":"B","logprob":-0.5},{"token":"C","logprob":-9.0}]}]}"#;
293
294 let dist = parse_distribution(shuffled, "m", Duration::ZERO).unwrap();
295
296 assert_eq!(dist.tokens[0].0, "B");
297 assert_eq!(dist.floor, -9.0);
298 }
299
300 #[test]
302 fn missing_logprobs_is_an_error_naming_the_model() {
303 let body = r#"{"model":"m","message":{"role":"assistant","content":"D"},"done":true}"#;
304
305 let err = parse_distribution(body, "tiny", Duration::ZERO).unwrap_err();
306
307 assert!(matches!(err, HostError::NoLogprobs { model } if model == "tiny"));
308 }
309
310 #[test]
311 fn empty_top_logprobs_is_treated_as_missing() {
312 let body = r#"{"logprobs":[{"token":"D","top_logprobs":[]}]}"#;
313
314 assert!(matches!(
315 parse_distribution(body, "m", Duration::ZERO),
316 Err(HostError::NoLogprobs { .. })
317 ));
318 }
319
320 #[test]
321 fn recognises_the_thinking_rejection() {
322 assert!(rejects_thinking(
323 r#"{"error":"registry.ollama.ai/library/x does not support thinking"}"#
324 ));
325 assert!(!rejects_thinking(r#"{"error":"model not found"}"#));
326 assert!(!rejects_thinking(
327 r#"{"error":"model \"deepthinker\" does not exist"}"#
328 ));
329 }
330
331 #[test]
334 fn required_sampling_options_are_serialised_verbatim() {
335 let json = serde_json::to_value(SamplingOptions::REQUIRED).unwrap();
336
337 assert_eq!(json["top_k"], 0);
338 assert_eq!(json["top_p"], 1.0);
339 assert_eq!(json["min_p"], 0.0);
340 assert_eq!(json["temperature"], 1.0);
341 assert_eq!(json["num_predict"], 1);
342 }
343
344 #[test]
345 fn request_body_disables_thinking_and_asks_for_logprobs() {
346 let host = OllamaHost::new("http://localhost:11434", Duration::from_secs(5)).unwrap();
347 let req = FirstTokenRequest {
348 model: "m".into(),
349 system: Some("sys".into()),
350 user: "usr".into(),
351 top_logprobs: 20,
352 keep_alive: Some("5m".into()),
353 };
354
355 let json = serde_json::to_value(host.body(&req, Some(false))).unwrap();
356
357 assert_eq!(json["think"], false);
358 assert_eq!(json["logprobs"], true);
359 assert_eq!(json["top_logprobs"], 20);
360 assert_eq!(json["stream"], false);
361 assert_eq!(json["messages"][0]["role"], "system");
362 assert_eq!(json["messages"][1]["content"], "usr");
363 }
364
365 #[test]
367 fn top_logprobs_is_clamped_to_the_host_ceiling() {
368 let host = OllamaHost::new("http://localhost:11434", Duration::from_secs(5)).unwrap();
369 let req = FirstTokenRequest {
370 model: "m".into(),
371 system: None,
372 user: "usr".into(),
373 top_logprobs: 99,
374 keep_alive: None,
375 };
376
377 let json = serde_json::to_value(host.body(&req, Some(false))).unwrap();
378
379 assert_eq!(json["top_logprobs"], 20);
380 assert!(json.get("keep_alive").is_none());
381 }
382
383 #[tokio::test]
386 async fn a_refused_think_flag_is_remembered_per_model() {
387 let mut server = mockito::Server::new_async().await;
388 let refused = server
389 .mock("POST", "/api/chat")
390 .match_body(mockito::Matcher::PartialJson(
391 serde_json::json!({"think": false}),
392 ))
393 .with_status(400)
394 .with_body(r#"{"error":"registry.ollama.ai/library/x does not support thinking"}"#)
395 .expect(1)
396 .create_async()
397 .await;
398 let answered = server
399 .mock("POST", "/api/chat")
400 .match_request(|request| {
404 !request
405 .utf8_lossy_body()
406 .is_ok_and(|body| body.contains(r#""think""#))
407 })
408 .with_body(REAL_RESPONSE)
409 .expect(2)
410 .create_async()
411 .await;
412 let host = OllamaHost::new(server.url(), Duration::from_secs(5)).unwrap();
413 let req = FirstTokenRequest {
414 model: "m".into(),
415 system: None,
416 user: "usr".into(),
417 top_logprobs: 20,
418 keep_alive: None,
419 };
420
421 host.first_token(req.clone()).await.unwrap();
422 host.first_token(req).await.unwrap();
423
424 refused.assert_async().await;
425 answered.assert_async().await;
426 }
427}