1mod ollama;
12mod openai;
13
14pub use ollama::OllamaHost;
15pub use openai::{Flavour, OpenAiCompatHost};
16
17use async_trait::async_trait;
18use std::sync::Arc;
19use std::time::Duration;
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum HostKind {
24 Ollama,
25 OpenAi,
27 Vllm,
28 LlamaCpp,
29 LmStudio,
30}
31
32impl HostKind {
33 pub fn default_url(self) -> &'static str {
35 match self {
36 HostKind::Ollama => "http://localhost:11434",
37 HostKind::OpenAi => "https://api.openai.com/v1",
38 HostKind::Vllm => "http://localhost:8000/v1",
39 HostKind::LlamaCpp => "http://localhost:8080/v1",
40 HostKind::LmStudio => "http://localhost:1234/v1",
41 }
42 }
43}
44
45#[derive(Debug, thiserror::Error)]
46#[error("unknown host {0:?}; expected one of ollama, openai, vllm, llamacpp, lmstudio")]
47pub struct UnknownHostKind(String);
48
49impl std::str::FromStr for HostKind {
50 type Err = UnknownHostKind;
51
52 fn from_str(s: &str) -> Result<Self, Self::Err> {
53 match s.trim().to_ascii_lowercase().as_str() {
54 "ollama" => Ok(HostKind::Ollama),
55 "openai" => Ok(HostKind::OpenAi),
56 "vllm" => Ok(HostKind::Vllm),
57 "llamacpp" | "llama.cpp" => Ok(HostKind::LlamaCpp),
58 "lmstudio" => Ok(HostKind::LmStudio),
59 _ => Err(UnknownHostKind(s.to_string())),
60 }
61 }
62}
63
64impl std::fmt::Display for HostKind {
65 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
66 f.write_str(match self {
67 HostKind::Ollama => "ollama",
68 HostKind::OpenAi => "openai",
69 HostKind::Vllm => "vllm",
70 HostKind::LlamaCpp => "llamacpp",
71 HostKind::LmStudio => "lmstudio",
72 })
73 }
74}
75
76pub fn connect(
78 kind: HostKind,
79 url: &str,
80 api_key: Option<String>,
81 timeout: Duration,
82) -> Result<Arc<dyn ModelHost>, HostError> {
83 let flavour = match kind {
84 HostKind::Ollama => return Ok(Arc::new(OllamaHost::new(url, timeout)?)),
85 HostKind::OpenAi => Flavour::Generic,
86 HostKind::Vllm => Flavour::Vllm,
87 HostKind::LlamaCpp => Flavour::LlamaCpp,
88 HostKind::LmStudio => Flavour::LmStudio,
89 };
90 Ok(Arc::new(OpenAiCompatHost::new(
91 url, api_key, flavour, timeout,
92 )?))
93}
94
95#[derive(Debug, Clone, Copy, PartialEq, Eq)]
97pub struct HostCapabilities {
98 pub max_top_logprobs: usize,
103}
104
105#[derive(Debug, Clone)]
107pub struct FirstTokenRequest {
108 pub model: String,
109 pub system: Option<String>,
110 pub user: String,
111 pub top_logprobs: usize,
112 pub keep_alive: Option<String>,
115}
116
117#[derive(Debug, Clone)]
119pub struct FirstTokenDistribution {
120 pub tokens: Vec<(String, f64)>,
122 pub floor: f64,
127 pub input_tokens: u32,
128 pub latency: Duration,
129}
130
131impl FirstTokenDistribution {
132 pub fn logprob(&self, token: &str) -> Option<f64> {
134 self.tokens
135 .iter()
136 .find(|(t, _)| t == token)
137 .map(|(_, lp)| *lp)
138 }
139}
140
141#[derive(Debug, thiserror::Error)]
142pub enum HostError {
143 #[error("host unreachable: {0}")]
144 Unavailable(String),
145
146 #[error("host timed out after {0:?}")]
147 Timeout(Duration),
148
149 #[error("host returned {status}: {body}")]
150 Status { status: u16, body: String },
151
152 #[error(
155 "host response carried no logprobs (model {model:?}); \
156 the runtime must support top_logprobs for cerno to work"
157 )]
158 NoLogprobs { model: String },
159
160 #[error("could not parse host response: {0}")]
161 Protocol(String),
162
163 #[error("host URL {url:?} is not usable: {reason}")]
166 InvalidUrl { url: String, reason: String },
167}
168
169fn base_url(raw: &str) -> Result<String, HostError> {
171 let trimmed = raw.trim().trim_end_matches('/');
172 let invalid = |reason: String| HostError::InvalidUrl {
173 url: raw.to_string(),
174 reason,
175 };
176
177 let parsed = reqwest::Url::parse(trimmed).map_err(|e| invalid(e.to_string()))?;
178 if !matches!(parsed.scheme(), "http" | "https") {
179 return Err(invalid("it must start with http:// or https://".into()));
182 }
183 if parsed.host_str().is_none() {
184 return Err(invalid("it names no host".into()));
185 }
186 Ok(trimmed.to_string())
187}
188
189const MAX_ERROR_BODY: usize = 300;
193
194impl HostError {
195 fn status(status: u16, body: String) -> Self {
201 tracing::warn!(status, body = %body, "host answered with an error");
202
203 if matches!(status, 401 | 403) {
204 return Self::Status {
205 status,
206 body: "the host refused cerno's credentials; the service log has its answer".into(),
207 };
208 }
209
210 let body = match body.char_indices().nth(MAX_ERROR_BODY) {
211 Some((cut, _)) => format!("{}…", &body[..cut]),
212 None => body,
213 };
214 Self::Status { status, body }
215 }
216
217 fn transport(error: &reqwest::Error, timeout: Duration) -> Self {
220 if error.is_timeout() {
221 Self::Timeout(timeout)
222 } else {
223 Self::Unavailable(error.to_string())
224 }
225 }
226}
227
228#[async_trait]
229pub trait ModelHost: Send + Sync {
230 fn capabilities(&self) -> HostCapabilities;
231
232 fn name(&self) -> &str;
234
235 async fn first_token(
236 &self,
237 req: FirstTokenRequest,
238 ) -> Result<FirstTokenDistribution, HostError>;
239}
240
241#[cfg(test)]
242mod tests {
243 use super::*;
244
245 #[test]
246 fn every_host_kind_parses_from_its_own_name() {
247 for kind in [
248 HostKind::Ollama,
249 HostKind::OpenAi,
250 HostKind::Vllm,
251 HostKind::LlamaCpp,
252 HostKind::LmStudio,
253 ] {
254 assert_eq!(kind.to_string().parse::<HostKind>().unwrap(), kind);
255 }
256 assert_eq!(
257 " LLAMA.CPP ".parse::<HostKind>().unwrap(),
258 HostKind::LlamaCpp
259 );
260 }
261
262 #[test]
263 fn an_unknown_host_kind_names_the_choices() {
264 let err = "tgi".parse::<HostKind>().unwrap_err().to_string();
265
266 assert!(err.contains("\"tgi\""), "{err}");
267 assert!(err.contains("lmstudio"), "{err}");
268 }
269
270 #[test]
272 fn default_urls_match_where_each_runtime_listens() {
273 assert_eq!(HostKind::Ollama.default_url(), "http://localhost:11434");
274 assert_eq!(HostKind::Vllm.default_url(), "http://localhost:8000/v1");
275 assert_eq!(HostKind::LlamaCpp.default_url(), "http://localhost:8080/v1");
276 assert_eq!(HostKind::LmStudio.default_url(), "http://localhost:1234/v1");
277 }
278
279 #[test]
280 fn a_long_error_body_is_cut_short() {
281 let HostError::Status { body, .. } = HostError::status(400, "x".repeat(5000)) else {
282 panic!()
283 };
284
285 assert_eq!(
286 body.chars().count(),
287 MAX_ERROR_BODY + 1,
288 "300 characters and an ellipsis"
289 );
290 assert!(body.ends_with('…'));
291 }
292
293 #[test]
294 fn an_authentication_failure_passes_on_none_of_the_body() {
295 for status in [401, 403] {
296 let HostError::Status { body, .. } = HostError::status(
297 status,
298 r#"{"error":{"message":"Incorrect API key provided: sk-proj-********abcd."}}"#
299 .into(),
300 ) else {
301 panic!()
302 };
303
304 assert!(!body.contains("sk-"), "{status}: {body}");
305 assert!(!body.contains("abcd"), "{status}: {body}");
306 }
307 }
308
309 #[test]
310 fn a_short_error_body_is_kept_whole() {
311 let HostError::Status { body, .. } = HostError::status(404, "model not found".into())
312 else {
313 panic!()
314 };
315
316 assert_eq!(body, "model not found");
317 }
318
319 #[test]
320 fn a_base_url_must_be_http_with_a_host() {
321 assert_eq!(
322 base_url(" http://localhost:11434/ ").unwrap(),
323 "http://localhost:11434"
324 );
325 assert_eq!(
326 base_url("https://api.openai.com/v1").unwrap(),
327 "https://api.openai.com/v1"
328 );
329
330 for bad in ["localhost:11434", "ftp://x", "http://", "not a url", ""] {
331 assert!(
332 matches!(base_url(bad), Err(HostError::InvalidUrl { .. })),
333 "{bad:?} was accepted"
334 );
335 }
336 }
337
338 #[test]
339 fn connect_refuses_a_url_without_a_scheme() {
340 let err = connect(
341 HostKind::Ollama,
342 "localhost:11434",
343 None,
344 Duration::from_secs(1),
345 )
346 .err()
347 .expect("refused");
348
349 assert!(err.to_string().contains("http://"), "{err}");
350 }
351
352 #[test]
353 fn connect_names_the_host_it_built() {
354 let timeout = Duration::from_secs(1);
355
356 assert_eq!(
357 connect(HostKind::Ollama, "http://x", None, timeout)
358 .unwrap()
359 .name(),
360 "ollama"
361 );
362 assert_eq!(
363 connect(HostKind::Vllm, "http://x/v1", None, timeout)
364 .unwrap()
365 .name(),
366 "vllm"
367 );
368 }
369}