Skip to main content

cerno_host/
ollama.rs

1//! Ollama adapter for [`ModelHost`].
2
3use 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
11/// Ollama reports at most this many ranked tokens per position; the server rejects more with
12/// "top_logprobs must be between 0 and 20".
13const OLLAMA_MAX_TOP_LOGPROBS: usize = 20;
14
15/// Sampling settings that make the reported distribution usable.
16///
17/// These are correctness conditions, not tuning knobs, and are deliberately not configurable:
18///
19/// * `num_predict: 1` — we want one token, the answer label. Nothing is generated after it.
20/// * `top_k: 0`, `top_p: 1`, `min_p: 0` — **the important one.** Ollama applies the sampling
21///   transform *before* reporting logprobs, so the default `top_k: 40` truncates the
22///   distribution and renormalises what survives. Measured against `gemma4:26b-a4b-it-q4_K_M`,
23///   that collapsed a four-option question to `A=0.0` with every rival at `-17`, and the
24///   remaining labels vanished from the list entirely. Disabling every filter restores the
25///   model's actual distribution.
26/// * `temperature: 1` — the identity for the same reason; calibration happens later, in
27///   `cerno-core`, where it is explicit and reversible.
28#[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    /// Suppresses the reasoning preamble. Without it the first generated token is a template
53    /// control token such as `<|channel|>` rather than the answer label.
54    #[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    /// Models that have refused `think` once. They are asked without it from then on, so the
74    /// refusal costs one extra round trip per model rather than one per question.
75    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
135/// Whether a failed response is Ollama complaining that the model has no thinking mode.
136///
137/// Models without a reasoning preamble reject `think` outright, so the same request has to go
138/// out again without the field. Only the message distinguishes this from a real failure, so it
139/// has to say both that something is unsupported and that it is thinking: a 404 for a model
140/// named `deepthinker` "does not exist", and a looser match would remember that model as unable
141/// to think and stop sending `think: false` to it for good.
142fn 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
194/// Pull the first position's ranked tokens out of an `/api/chat` response.
195fn 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    // `top_logprobs` holds the ranked alternatives and always includes the sampled token itself.
212    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    /// A real `/api/chat` response, trimmed. Captured from `gemma4:26b-a4b-it-q4_K_M`.
259    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        // The floor is the weakest reported entry, the bound for anything not listed.
286        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    /// A response with no logprobs is unusable, and saying so plainly beats a silent zero.
301    #[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    /// The sampling options are the whole reason the distribution is readable; pin them so a
332    /// well-meaning edit cannot quietly reintroduce the default `top_k`.
333    #[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    /// Asking for more than Ollama allows would be rejected outright, so clamp instead.
366    #[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    /// A model without a thinking mode refuses `think`. The first question pays for one retry;
384    /// every later question to that model goes out without the field straight away.
385    #[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            // Only a body without the switch is answered. Without this guard mockito would hand a
401            // repeated switch to this mock once the refusal had been used up, and the test
402            // would pass whether or not the refusal was remembered.
403            .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}