1use 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
14const OPENAI_MAX_TOP_LOGPROBS: usize = 20;
17
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20pub enum Flavour {
21 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 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 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 refused_thinking: Mutex<HashSet<String>>,
83}
84
85impl OpenAiCompatHost {
86 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
161fn 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
217fn 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 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 #[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 assert!(json.get("keep_alive").is_none(), "{flavour:?}");
382 }
383 }
384
385 #[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 #[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 #[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 #[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 .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 host.first_token(request(4)).await.unwrap();
538
539 refused.assert_async().await;
540 answered.assert_async().await;
541 }
542}