Skip to main content

cerno_host/
openai.rs

1//! Adapter for anything that speaks OpenAI's `/v1/chat/completions`: vLLM, llama.cpp's
2//! `llama-server`, LM Studio, OpenAI itself.
3//!
4//! The response shape is the same everywhere, so there is one parser. What differs is which
5//! sampling fields a runtime accepts, and that is all a [`Flavour`] carries.
6
7use crate::{FirstTokenDistribution, FirstTokenRequest, HostCapabilities, HostError, ModelHost};
8use async_trait::async_trait;
9use serde_json::{Map, Value, json};
10use std::collections::HashSet;
11use std::sync::Mutex;
12use std::time::{Duration, Instant};
13
14/// OpenAI rejects more than this, and it is vLLM's default `--max-logprobs`. Past 20 there are
15/// no labels to observe anyway.
16const OPENAI_MAX_TOP_LOGPROBS: usize = 20;
17
18/// Which OpenAI-compatible runtime is on the other end.
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20pub enum Flavour {
21    /// Standard fields only. OpenAI answers 400 to a field it does not know, so nothing
22    /// runtime-specific can be sent — `top_k` and `min_p` included.
23    ///
24    /// Measured behind Ollama's `/v1` on `gemma4:26b-a4b-it-q4_K_M`: this body returns the same
25    /// distribution as the native path with its pinned options, to the third decimal.
26    Generic,
27    Vllm,
28    LlamaCpp,
29    LmStudio,
30}
31
32impl Flavour {
33    fn name(self) -> &'static str {
34        match self {
35            Flavour::Generic => "openai",
36            Flavour::Vllm => "vllm",
37            Flavour::LlamaCpp => "llamacpp",
38            Flavour::LmStudio => "lmstudio",
39        }
40    }
41
42    /// Sampling settings on top of the standard ones, for the same reason as Ollama's
43    /// `SamplingOptions::REQUIRED`: a runtime that applies its sampling transform before reporting
44    /// logprobs hands back a truncated, renormalised distribution unless every filter is off.
45    ///
46    /// * vLLM reports raw logprobs by default, but pinning `top_k: -1` (its "disabled") costs
47    ///   nothing and survives a server started with a different `--logprobs-mode`.
48    /// * llama.cpp disables `top_k` with `0`, and `post_sampling_probs: false` asks for the
49    ///   distribution before the sampler chain.
50    fn sampling(self) -> Value {
51        match self {
52            Flavour::Generic => json!({}),
53            Flavour::Vllm => json!({"top_k": -1, "min_p": 0.0}),
54            Flavour::LlamaCpp => json!({"top_k": 0, "min_p": 0.0, "post_sampling_probs": false}),
55            Flavour::LmStudio => json!({"top_k": 0, "min_p": 0.0}),
56        }
57    }
58
59    /// Fields that switch the reasoning preamble off. Without them the first token is a template
60    /// control token — measured on `gemma4:e2b-it-qat` behind Ollama's `/v1`: `<|channel>` at
61    /// `-0.03`, the answer at `-3.5`. `reasoning_effort: "none"` is the standard switch and fixed
62    /// it there; `chat_template_kwargs` is how vLLM and llama.cpp reach Qwen-style templates.
63    fn no_thinking(self) -> Value {
64        match self {
65            Flavour::Vllm | Flavour::LlamaCpp => json!({
66                "reasoning_effort": "none",
67                "chat_template_kwargs": {"enable_thinking": false},
68            }),
69            Flavour::Generic | Flavour::LmStudio => json!({"reasoning_effort": "none"}),
70        }
71    }
72}
73
74pub struct OpenAiCompatHost {
75    client: reqwest::Client,
76    base_url: String,
77    api_key: Option<String>,
78    flavour: Flavour,
79    timeout: Duration,
80    /// Models that have refused the thinking switch once. They are asked without it from then
81    /// on, so the refusal costs one extra round trip per model rather than one per question.
82    refused_thinking: Mutex<HashSet<String>>,
83}
84
85impl OpenAiCompatHost {
86    /// `base_url` includes the version prefix, e.g. `http://localhost:8000/v1`.
87    pub fn new(
88        base_url: impl Into<String>,
89        api_key: Option<String>,
90        flavour: Flavour,
91        timeout: Duration,
92    ) -> Result<Self, HostError> {
93        let client = reqwest::Client::builder()
94            .timeout(timeout)
95            .build()
96            .map_err(|e| HostError::Unavailable(e.to_string()))?;
97        Ok(Self {
98            client,
99            base_url: crate::base_url(&base_url.into())?,
100            api_key,
101            flavour,
102            timeout,
103            refused_thinking: Mutex::default(),
104        })
105    }
106
107    fn body(&self, req: &FirstTokenRequest, think_off: bool) -> Value {
108        let mut messages = Vec::with_capacity(2);
109        if let Some(system) = req.system.as_deref() {
110            messages.push(json!({"role": "system", "content": system}));
111        }
112        messages.push(json!({"role": "user", "content": req.user}));
113
114        let mut body = json!({
115            "model": req.model,
116            "messages": messages,
117            "stream": false,
118            "logprobs": true,
119            "top_logprobs": req.top_logprobs.min(OPENAI_MAX_TOP_LOGPROBS),
120            "max_tokens": 1,
121            "temperature": 1.0,
122            "top_p": 1.0,
123        });
124        let object = body.as_object_mut().expect("body is an object");
125        merge(object, self.flavour.sampling());
126        if think_off {
127            merge(object, self.flavour.no_thinking());
128        }
129        body
130    }
131
132    async fn post(&self, body: &Value) -> Result<(u16, String), HostError> {
133        let mut request = self
134            .client
135            .post(format!("{}/chat/completions", self.base_url))
136            .json(body);
137        if let Some(key) = self.api_key.as_deref() {
138            request = request.bearer_auth(key);
139        }
140
141        let response = request
142            .send()
143            .await
144            .map_err(|e| HostError::transport(&e, self.timeout))?;
145
146        let status = response.status().as_u16();
147        let text = response
148            .text()
149            .await
150            .map_err(|e| HostError::transport(&e, self.timeout))?;
151        Ok((status, text))
152    }
153}
154
155fn merge(into: &mut Map<String, Value>, extra: Value) {
156    if let Value::Object(extra) = extra {
157        into.extend(extra);
158    }
159}
160
161/// Whether a failed response is the runtime refusing the thinking switch — OpenAI does for any
162/// model without a reasoning mode. The request then goes out again without it.
163fn rejects_thinking(body: &str) -> bool {
164    let lower = body.to_ascii_lowercase();
165    lower.contains("reasoning_effort")
166        || lower.contains("chat_template_kwargs")
167        || lower.contains("enable_thinking")
168}
169
170#[async_trait]
171impl ModelHost for OpenAiCompatHost {
172    fn capabilities(&self) -> HostCapabilities {
173        HostCapabilities {
174            max_top_logprobs: OPENAI_MAX_TOP_LOGPROBS,
175        }
176    }
177
178    fn name(&self) -> &str {
179        self.flavour.name()
180    }
181
182    async fn first_token(
183        &self,
184        req: FirstTokenRequest,
185    ) -> Result<FirstTokenDistribution, HostError> {
186        let started = Instant::now();
187
188        let think_off = !self
189            .refused_thinking
190            .lock()
191            .expect("lock is never poisoned")
192            .contains(&req.model);
193
194        let (mut status, mut text) = self.post(&self.body(&req, think_off)).await?;
195
196        if think_off && (400..500).contains(&status) && rejects_thinking(&text) {
197            tracing::debug!(
198                model = %req.model,
199                host = self.flavour.name(),
200                "host rejects the thinking switch; retrying without it, and from now on"
201            );
202            self.refused_thinking
203                .lock()
204                .expect("lock is never poisoned")
205                .insert(req.model.clone());
206            (status, text) = self.post(&self.body(&req, false)).await?;
207        }
208
209        if status >= 400 {
210            return Err(HostError::status(status, text));
211        }
212
213        parse_distribution(&text, &req.model, started.elapsed())
214    }
215}
216
217/// Pull the first position's ranked tokens out of a `/chat/completions` response.
218fn parse_distribution(
219    text: &str,
220    model: &str,
221    latency: Duration,
222) -> Result<FirstTokenDistribution, HostError> {
223    let value: Value =
224        serde_json::from_str(text).map_err(|e| HostError::Protocol(e.to_string()))?;
225    let no_logprobs = || HostError::NoLogprobs {
226        model: model.to_string(),
227    };
228
229    let ranked = value
230        .pointer("/choices/0/logprobs/content/0/top_logprobs")
231        .and_then(Value::as_array)
232        .ok_or_else(no_logprobs)?;
233
234    let mut tokens: Vec<(String, f64)> = ranked
235        .iter()
236        .filter_map(|entry| {
237            let token = entry.get("token")?.as_str()?.to_string();
238            let logprob = entry.get("logprob")?.as_f64()?;
239            Some((token, logprob))
240        })
241        .collect();
242
243    if tokens.is_empty() {
244        return Err(no_logprobs());
245    }
246
247    tokens.sort_by(|a, b| b.1.total_cmp(&a.1));
248
249    let floor = tokens
250        .last()
251        .map(|(_, lp)| *lp)
252        .expect("tokens is non-empty");
253
254    let input_tokens = value
255        .pointer("/usage/prompt_tokens")
256        .and_then(Value::as_u64)
257        .unwrap_or(0) as u32;
258
259    Ok(FirstTokenDistribution {
260        tokens,
261        floor,
262        input_tokens,
263        latency,
264    })
265}
266
267#[cfg(test)]
268mod tests {
269    use super::*;
270
271    /// A real `/v1/chat/completions` response, trimmed of `bytes`. Captured from
272    /// `gemma4:e2b-it-qat` behind Ollama's OpenAI endpoint, with `reasoning_effort: "none"`.
273    const REAL_RESPONSE: &str = r#"{
274        "id": "chatcmpl-147",
275        "object": "chat.completion",
276        "model": "gemma4:e2b-it-qat",
277        "choices": [{
278            "index": 0,
279            "message": {"role": "assistant", "content": "B"},
280            "finish_reason": "length",
281            "logprobs": {"content": [{
282                "token": "B",
283                "logprob": -0.002,
284                "top_logprobs": [
285                    {"token": "B", "logprob": -0.002},
286                    {"token": "b", "logprob": -7.721},
287                    {"token": "**", "logprob": -7.946},
288                    {"token": "", "logprob": -7.976},
289                    {"token": "C", "logprob": -8.674}
290                ]
291            }]}
292        }],
293        "usage": {"prompt_tokens": 40, "completion_tokens": 1, "total_tokens": 41}
294    }"#;
295
296    fn request(top_logprobs: usize) -> FirstTokenRequest {
297        FirstTokenRequest {
298            model: "m".into(),
299            system: Some("sys".into()),
300            user: "usr".into(),
301            top_logprobs,
302            keep_alive: Some("5m".into()),
303        }
304    }
305
306    fn body(flavour: Flavour, think_off: bool) -> Value {
307        OpenAiCompatHost::new(
308            "http://localhost:8000/v1",
309            None,
310            flavour,
311            Duration::from_secs(5),
312        )
313        .unwrap()
314        .body(&request(20), think_off)
315    }
316
317    #[test]
318    fn parses_ranked_tokens_and_floor() {
319        let dist = parse_distribution(REAL_RESPONSE, "m", Duration::from_millis(7)).unwrap();
320
321        assert_eq!(dist.tokens.len(), 5);
322        assert_eq!(dist.tokens[0].0, "B");
323        assert_eq!(dist.input_tokens, 40);
324        assert_eq!(dist.logprob("C"), Some(-8.674));
325        assert_eq!(dist.logprob("A"), None);
326        assert_eq!(dist.floor, -8.674);
327    }
328
329    #[test]
330    fn ranks_tokens_even_when_the_host_reports_them_unsorted() {
331        let shuffled = r#"{"choices":[{"logprobs":{"content":[{"token":"B","top_logprobs":[
332            {"token":"A","logprob":-5.0},{"token":"B","logprob":-0.5},{"token":"C","logprob":-9.0}]}]}}]}"#;
333
334        let dist = parse_distribution(shuffled, "m", Duration::ZERO).unwrap();
335
336        assert_eq!(dist.tokens[0].0, "B");
337        assert_eq!(dist.floor, -9.0);
338    }
339
340    /// A runtime that ignores `logprobs` answers normally with `"logprobs": null`.
341    #[test]
342    fn missing_logprobs_is_an_error_naming_the_model() {
343        let body =
344            r#"{"choices":[{"message":{"role":"assistant","content":"B"},"logprobs":null}]}"#;
345
346        let err = parse_distribution(body, "tiny", Duration::ZERO).unwrap_err();
347
348        assert!(matches!(err, HostError::NoLogprobs { model } if model == "tiny"));
349    }
350
351    #[test]
352    fn empty_top_logprobs_is_treated_as_missing() {
353        let body = r#"{"choices":[{"logprobs":{"content":[{"token":"B","top_logprobs":[]}]}}]}"#;
354
355        assert!(matches!(
356            parse_distribution(body, "m", Duration::ZERO),
357            Err(HostError::NoLogprobs { .. })
358        ));
359    }
360
361    #[test]
362    fn the_standard_fields_are_the_same_for_every_flavour() {
363        for flavour in [
364            Flavour::Generic,
365            Flavour::Vllm,
366            Flavour::LlamaCpp,
367            Flavour::LmStudio,
368        ] {
369            let json = body(flavour, true);
370
371            assert_eq!(json["logprobs"], true, "{flavour:?}");
372            assert_eq!(json["top_logprobs"], 20, "{flavour:?}");
373            assert_eq!(json["max_tokens"], 1, "{flavour:?}");
374            assert_eq!(json["temperature"], 1.0, "{flavour:?}");
375            assert_eq!(json["top_p"], 1.0, "{flavour:?}");
376            assert_eq!(json["stream"], false, "{flavour:?}");
377            assert_eq!(json["reasoning_effort"], "none", "{flavour:?}");
378            assert_eq!(json["messages"][0]["role"], "system");
379            assert_eq!(json["messages"][1]["content"], "usr");
380            // keep_alive is Ollama's; nobody else knows it.
381            assert!(json.get("keep_alive").is_none(), "{flavour:?}");
382        }
383    }
384
385    /// OpenAI answers 400 to a field it does not know, so the generic body carries none.
386    #[test]
387    fn the_generic_flavour_sends_nothing_nonstandard() {
388        let json = body(Flavour::Generic, true);
389
390        for field in [
391            "top_k",
392            "min_p",
393            "post_sampling_probs",
394            "chat_template_kwargs",
395        ] {
396            assert!(json.get(field).is_none(), "{field}");
397        }
398    }
399
400    /// The sampling fields are the reason the distribution is readable; pin them verbatim.
401    #[test]
402    fn runtime_sampling_fields_are_serialised_verbatim() {
403        let vllm = body(Flavour::Vllm, true);
404        assert_eq!(vllm["top_k"], -1);
405        assert_eq!(vllm["min_p"], 0.0);
406        assert_eq!(vllm["chat_template_kwargs"]["enable_thinking"], false);
407
408        let llamacpp = body(Flavour::LlamaCpp, true);
409        assert_eq!(llamacpp["top_k"], 0);
410        assert_eq!(llamacpp["min_p"], 0.0);
411        assert_eq!(llamacpp["post_sampling_probs"], false);
412        assert_eq!(llamacpp["chat_template_kwargs"]["enable_thinking"], false);
413
414        let lmstudio = body(Flavour::LmStudio, true);
415        assert_eq!(lmstudio["top_k"], 0);
416        assert_eq!(lmstudio["min_p"], 0.0);
417    }
418
419    /// The retry drops the thinking switch and nothing else — least of all the sampling fields.
420    #[test]
421    fn the_retry_body_keeps_sampling_and_drops_only_the_thinking_switch() {
422        let json = body(Flavour::Vllm, false);
423
424        assert!(json.get("reasoning_effort").is_none());
425        assert!(json.get("chat_template_kwargs").is_none());
426        assert_eq!(json["top_k"], -1);
427    }
428
429    #[test]
430    fn top_logprobs_is_clamped_to_the_host_ceiling() {
431        let host = OpenAiCompatHost::new(
432            "http://x/v1",
433            None,
434            Flavour::Generic,
435            Duration::from_secs(5),
436        )
437        .unwrap();
438
439        assert_eq!(host.body(&request(99), true)["top_logprobs"], 20);
440    }
441
442    #[test]
443    fn recognises_the_thinking_rejection() {
444        assert!(rejects_thinking(
445            r#"{"error":{"message":"Unsupported parameter: 'reasoning_effort' is not supported with this model.","param":"reasoning_effort"}}"#
446        ));
447        assert!(!rejects_thinking(
448            r#"{"error":{"message":"model not found"}}"#
449        ));
450    }
451
452    #[tokio::test]
453    async fn the_api_key_is_sent_as_a_bearer_token() {
454        let mut server = mockito::Server::new_async().await;
455        let mock = server
456            .mock("POST", "/v1/chat/completions")
457            .match_header("authorization", "Bearer sk-test")
458            .with_body(REAL_RESPONSE)
459            .create_async()
460            .await;
461        let host = OpenAiCompatHost::new(
462            format!("{}/v1/", server.url()),
463            Some("sk-test".into()),
464            Flavour::Generic,
465            Duration::from_secs(5),
466        )
467        .unwrap();
468
469        host.first_token(request(4)).await.unwrap();
470
471        mock.assert_async().await;
472    }
473
474    #[tokio::test]
475    async fn without_a_key_no_authorization_header_is_sent() {
476        let mut server = mockito::Server::new_async().await;
477        let mock = server
478            .mock("POST", "/v1/chat/completions")
479            .match_header("authorization", mockito::Matcher::Missing)
480            .with_body(REAL_RESPONSE)
481            .create_async()
482            .await;
483        let host = OpenAiCompatHost::new(
484            format!("{}/v1", server.url()),
485            None,
486            Flavour::LlamaCpp,
487            Duration::from_secs(5),
488        )
489        .unwrap();
490
491        host.first_token(request(4)).await.unwrap();
492
493        mock.assert_async().await;
494    }
495
496    /// A model without a reasoning mode refuses the switch; the same question goes out again
497    /// without it and is answered, and later questions to that model skip the refusal.
498    #[tokio::test]
499    async fn a_rejected_thinking_switch_is_retried_without_it_and_remembered() {
500        let mut server = mockito::Server::new_async().await;
501        let refused = server
502            .mock("POST", "/v1/chat/completions")
503            .match_body(mockito::Matcher::PartialJson(
504                json!({"reasoning_effort": "none"}),
505            ))
506            .with_status(400)
507            .with_body(r#"{"error":{"message":"Unsupported parameter: 'reasoning_effort'"}}"#)
508            .expect(1)
509            .create_async()
510            .await;
511        let answered = server
512            .mock("POST", "/v1/chat/completions")
513            // Only a body without the switch is answered. Without this guard mockito would hand a
514            // repeated switch to this mock once the refusal had been used up, and the test
515            // would pass whether or not the refusal was remembered.
516            .match_request(|request| {
517                !request
518                    .utf8_lossy_body()
519                    .is_ok_and(|body| body.contains(r#""reasoning_effort""#))
520            })
521            .with_body(REAL_RESPONSE)
522            .expect(2)
523            .create_async()
524            .await;
525        let host = OpenAiCompatHost::new(
526            format!("{}/v1", server.url()),
527            None,
528            Flavour::Generic,
529            Duration::from_secs(5),
530        )
531        .unwrap();
532
533        let dist = host.first_token(request(4)).await.unwrap();
534        assert_eq!(dist.tokens[0].0, "B");
535
536        // The second question to the same model must not pay for the refusal again.
537        host.first_token(request(4)).await.unwrap();
538
539        refused.assert_async().await;
540        answered.assert_async().await;
541    }
542}