Skip to main content

cerno_host/
lib.rs

1//! The seam between cerno and whatever runs the model.
2//!
3//! A host knows nothing about Noul, Choice or Score. It answers exactly one question: given a
4//! prompt, what is the probability distribution over the *first* token the model would generate?
5//! All primitive logic lives in `cerno-core` and therefore holds for every host equally.
6//!
7//! Two adapters cover the runtimes: [`OllamaHost`] for Ollama's native API, and
8//! [`OpenAiCompatHost`] for everything speaking `/v1/chat/completions`. [`connect`] picks one
9//! from a [`HostKind`].
10
11mod 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/// Which runtime a server talks to. One per process.
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum HostKind {
24    Ollama,
25    /// OpenAI itself, or any server that speaks its API and nothing more.
26    OpenAi,
27    Vllm,
28    LlamaCpp,
29    LmStudio,
30}
31
32impl HostKind {
33    /// Where the runtime listens out of the box.
34    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
76/// Build the adapter for `kind`. Ollama takes no API key; the others send it as a bearer token.
77pub 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/// What a host can and cannot do. `cerno-core` reads this to size its label alphabet.
96#[derive(Debug, Clone, Copy, PartialEq, Eq)]
97pub struct HostCapabilities {
98    /// How many ranked tokens the host will report for one position.
99    ///
100    /// Ollama and OpenAI cap this at 20. A label beyond that rank is invisible, which is exactly
101    /// the bound that limits a Choice to 20 options.
102    pub max_top_logprobs: usize,
103}
104
105/// One prompt, one token, full distribution.
106#[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    /// How long the host should keep the model resident after answering. Only Ollama has this;
113    /// the OpenAI-compatible runtimes keep their model loaded for the life of the process.
114    pub keep_alive: Option<String>,
115}
116
117/// The distribution over the first generated token.
118#[derive(Debug, Clone)]
119pub struct FirstTokenDistribution {
120    /// Ranked tokens, highest logprob first.
121    pub tokens: Vec<(String, f64)>,
122    /// The lowest logprob the host reported.
123    ///
124    /// A label missing from `tokens` ranked below every entry, so this value is a strict upper
125    /// bound on its logprob. `cerno-core` substitutes it and flags the answer as truncated.
126    pub floor: f64,
127    pub input_tokens: u32,
128    pub latency: Duration,
129}
130
131impl FirstTokenDistribution {
132    /// The logprob of `token`, or `None` when it fell outside the reported window.
133    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    /// The host answered, but not in the shape we need — most often because the model or the
153    /// runtime does not report logprobs at all.
154    #[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    /// A base URL no request could be sent to. Caught when the host is built, because otherwise
164    /// the server starts and every request fails with reqwest's "builder error".
165    #[error("host URL {url:?} is not usable: {reason}")]
166    InvalidUrl { url: String, reason: String },
167}
168
169/// The base URL with any trailing slash removed, or why it cannot be one.
170fn 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        // `localhost:11434` parses, with `localhost` as its scheme, so naming the scheme found
180        // would only confuse.
181        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
189/// How much of a failed response's body travels on in the error. The whole body is logged; the
190/// error reaches the caller of the service, and an upstream body can carry account or request
191/// details that are the operator's business, not theirs.
192const MAX_ERROR_BODY: usize = 300;
193
194impl HostError {
195    /// A non-success status from the host, with its body logged in full and kept short.
196    ///
197    /// An authentication failure keeps none of it. OpenAI answers a wrong key with "Incorrect
198    /// API key provided: sk-…abcd", and part of the operator's key is not the caller's business
199    /// however short the excerpt.
200    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    /// Map a transport failure, keeping a timeout a timeout wherever in the exchange it struck —
218    /// sending the request or reading the body.
219    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    /// A human-readable name for the host, used in logs and diagnostics.
233    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    /// The OpenAI-compatible defaults carry the `/v1` prefix; Ollama's native API has none.
271    #[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}