Skip to main content

cerno_core/
engine.rs

1//! The one path all three primitives take.
2//!
3//! Noul, Choice and Score differ only in what the options are and how the resulting
4//! distribution is named. Everything between — labelling, prompting, the single host call,
5//! folding token variants, substituting the floor, calibrating, normalising — is shared.
6
7use crate::{labels, math, prompt};
8use cerno_host::{FirstTokenDistribution, FirstTokenRequest, HostError, ModelHost};
9use cerno_types::{
10    Answer, Calibration, ErrorCode, LevelProbability, MAX_LEVELS, MAX_OPTIONS, MAX_QUESTIONS,
11    MIN_LEVELS, OptionProbability, Question, QuestionKind, SystemOneRequest,
12};
13use std::collections::BTreeMap;
14use std::sync::Arc;
15
16#[derive(Debug, thiserror::Error)]
17pub enum EngineError {
18    #[error("{message}")]
19    Invalid {
20        code: ErrorCode,
21        message: String,
22        question_id: Option<String>,
23    },
24
25    /// The model answered, but with nothing resembling one of the labels we offered.
26    ///
27    /// Distinct from a host failure: the runtime worked, the model simply did not follow the
28    /// instruction. Almost always this means the model is too small or too chatty for the job,
29    /// which is what the benchmark in `cerno-bench` exists to catch before deployment.
30    #[error(
31        "model answered question {question_id:?} with none of the labels {expected:?}; \
32         it offered {observed:?} instead"
33    )]
34    NoLabelMatched {
35        question_id: String,
36        expected: Vec<String>,
37        observed: Vec<String>,
38    },
39
40    /// The host answered 404 to a question naming `model`.
41    ///
42    /// Every runtime answers a model it does not have that way, and without `strict_models`
43    /// that is where a misspelt `model` surfaces — the caller's to fix, so not a 5xx inviting a
44    /// retry. But a host URL missing its `/v1` answers 404 to every model as well, so the
45    /// message says what to check when it is not the model.
46    #[error(
47        "model {model:?} is not known to {host} ({source}); if every model is refused this way, \
48         the host URL is probably wrong"
49    )]
50    UnknownModel {
51        model: String,
52        host: String,
53        source: HostError,
54    },
55
56    #[error(transparent)]
57    Host(#[from] HostError),
58}
59
60impl EngineError {
61    pub fn code(&self) -> ErrorCode {
62        match self {
63            Self::Invalid { code, .. } => *code,
64            Self::NoLabelMatched { .. } => ErrorCode::NoLabelMatched,
65            Self::UnknownModel { .. } => ErrorCode::UnknownModel,
66            Self::Host(HostError::Timeout(_)) => ErrorCode::HostTimeout,
67            Self::Host(_) => ErrorCode::HostUnavailable,
68        }
69    }
70
71    pub fn question_id(&self) -> Option<&str> {
72        match self {
73            Self::Invalid { question_id, .. } => question_id.as_deref(),
74            Self::NoLabelMatched { question_id, .. } => Some(question_id),
75            Self::UnknownModel { .. } | Self::Host(_) => None,
76        }
77    }
78
79    fn invalid(code: ErrorCode, question_id: Option<&str>, message: impl Into<String>) -> Self {
80        Self::Invalid {
81            code,
82            message: message.into(),
83            question_id: question_id.map(str::to_string),
84        }
85    }
86}
87
88/// What a question looks like once the primitive has been stripped away.
89struct Ballot<'a> {
90    question: Option<&'a str>,
91    /// The option texts shown to the model, in request order.
92    options: Vec<String>,
93}
94
95/// One question's resolved distribution over its labels.
96struct Tally {
97    /// Probability per option, in request order.
98    probabilities: Vec<f64>,
99    /// The logprob actually fed into the softmax, per label. A floor substitution appears here
100    /// like any other value; `truncated` is what says one happened.
101    logprobs: BTreeMap<String, f64>,
102    confidence: f64,
103    /// The labels that fell outside the reported window and were given the floor.
104    truncated_labels: Vec<String>,
105    /// The probability the model itself put on the observed labels, before any normalising.
106    label_mass: f64,
107}
108
109pub struct Engine {
110    host: Arc<dyn ModelHost>,
111    keep_alive: Option<String>,
112}
113
114impl Engine {
115    pub fn new(host: Arc<dyn ModelHost>, keep_alive: Option<String>) -> Self {
116        Self { host, keep_alive }
117    }
118
119    /// The largest option count this engine can answer, given what its host reports.
120    pub fn max_options(&self) -> usize {
121        MAX_OPTIONS.min(self.host.capabilities().max_top_logprobs)
122    }
123
124    pub fn host_name(&self) -> &str {
125        self.host.name()
126    }
127
128    /// Reject a request that cannot be answered, before any model is loaded.
129    pub fn validate(&self, request: &SystemOneRequest) -> Result<(), EngineError> {
130        if request.state.trim().is_empty() {
131            return Err(EngineError::invalid(
132                ErrorCode::EmptyState,
133                None,
134                "state must not be empty",
135            ));
136        }
137
138        if request.questions.is_empty() {
139            return Err(EngineError::invalid(
140                ErrorCode::NoQuestions,
141                None,
142                "at least one question is required",
143            ));
144        }
145
146        if let Some(calibration) = request.calibration {
147            let t = calibration.temperature;
148            if !t.is_finite() || t <= 0.0 {
149                return Err(EngineError::invalid(
150                    ErrorCode::InvalidCalibration,
151                    None,
152                    format!("calibration temperature must be finite and above zero, got {t}"),
153                ));
154            }
155        }
156
157        if request.questions.len() > MAX_QUESTIONS {
158            return Err(EngineError::invalid(
159                ErrorCode::TooManyQuestions,
160                None,
161                format!(
162                    "a request may have at most {MAX_QUESTIONS} questions, got {}; send the rest \
163                     in a second request",
164                    request.questions.len()
165                ),
166            ));
167        }
168
169        let mut seen = BTreeMap::new();
170        for question in &request.questions {
171            if question.id.trim().is_empty() {
172                return Err(EngineError::invalid(
173                    ErrorCode::EmptyQuestionId,
174                    Some(&question.id),
175                    "question ids must not be empty; the answer is found by its id",
176                ));
177            }
178            if seen.insert(question.id.as_str(), ()).is_some() {
179                return Err(EngineError::invalid(
180                    ErrorCode::DuplicateQuestionId,
181                    Some(&question.id),
182                    format!("question id {:?} appears more than once", question.id),
183                ));
184            }
185            self.validate_question(question)?;
186        }
187
188        Ok(())
189    }
190
191    fn validate_question(&self, question: &Question) -> Result<(), EngineError> {
192        let id = Some(question.id.as_str());
193        let max = self.max_options();
194
195        match &question.kind {
196            QuestionKind::Noul(text) => {
197                if text.trim().is_empty() {
198                    return Err(EngineError::invalid(
199                        ErrorCode::EmptyQuestion,
200                        id,
201                        "a noul question must not be empty",
202                    ));
203                }
204            }
205            QuestionKind::Choice(spec) => {
206                if spec.options.len() < 2 {
207                    return Err(EngineError::invalid(
208                        ErrorCode::TooFewOptions,
209                        id,
210                        format!(
211                            "a choice needs at least 2 options, got {}",
212                            spec.options.len()
213                        ),
214                    ));
215                }
216                if spec.options.len() > max {
217                    return Err(EngineError::invalid(
218                        ErrorCode::TooManyOptions,
219                        id,
220                        format!(
221                            "a choice may have at most {max} options, got {}; the host reports \
222                             only {max} ranked tokens, so further options cannot be observed. \
223                             Split them across a first question that picks a group and a second \
224                             that picks within it.",
225                            spec.options.len()
226                        ),
227                    ));
228                }
229                if spec.options.iter().any(|o| o.trim().is_empty()) {
230                    return Err(EngineError::invalid(
231                        ErrorCode::EmptyQuestion,
232                        id,
233                        "choice options must not be empty",
234                    ));
235                }
236                if let Some(twice) = repeated(&spec.options) {
237                    return Err(EngineError::invalid(
238                        ErrorCode::DuplicateOption,
239                        id,
240                        format!(
241                            "option {twice:?} appears more than once; the model would see it \
242                             twice and split its probability between the copies"
243                        ),
244                    ));
245                }
246            }
247            QuestionKind::Score(spec) => {
248                let count = spec.levels.count();
249                if count < MIN_LEVELS as usize || count > MAX_LEVELS as usize {
250                    return Err(EngineError::invalid(
251                        ErrorCode::InvalidLevels,
252                        id,
253                        format!(
254                            "a score rubric needs between {MIN_LEVELS} and {MAX_LEVELS} levels, \
255                             got {count}"
256                        ),
257                    ));
258                }
259                let legend = spec.levels.legend();
260                if legend.iter().any(|l| l.trim().is_empty()) {
261                    return Err(EngineError::invalid(
262                        ErrorCode::EmptyQuestion,
263                        id,
264                        "score level labels must not be empty",
265                    ));
266                }
267                if let Some(twice) = repeated(&legend) {
268                    return Err(EngineError::invalid(
269                        ErrorCode::DuplicateOption,
270                        id,
271                        format!(
272                            "level {twice:?} appears more than once; the model could not tell \
273                             the two apart"
274                        ),
275                    ));
276                }
277            }
278        }
279
280        Ok(())
281    }
282
283    /// Ask one question and type its answer. The `u32` is the prompt's token count.
284    pub async fn answer(
285        &self,
286        state: &str,
287        question: &Question,
288        model: &str,
289        calibration: Calibration,
290    ) -> Result<(Answer, u32), EngineError> {
291        let (answer, distribution) = self
292            .answer_with_distribution(state, question, model, calibration)
293            .await?;
294        Ok((answer, distribution.input_tokens))
295    }
296
297    /// [`Engine::answer`], also handing back the distribution the answer was read from.
298    ///
299    /// The answer renormalises over the offered labels, so it cannot say whether the model's
300    /// most likely token was a label at all. The distribution can, which is what the benchmark
301    /// measures fidelity by.
302    pub async fn answer_with_distribution(
303        &self,
304        state: &str,
305        question: &Question,
306        model: &str,
307        calibration: Calibration,
308    ) -> Result<(Answer, FirstTokenDistribution), EngineError> {
309        self.validate_question(question)?;
310
311        let ballot = ballot_for(&question.kind);
312        let (tally, distribution) = self
313            .tally(state, &ballot, model, calibration, &question.id)
314            .await?;
315
316        let answer = match &question.kind {
317            QuestionKind::Noul(_) => Answer::Noul {
318                // Option 0 is "Yes" by construction; see `ballot_for`.
319                noul: tally.probabilities[0],
320                raw_logprobs: tally.logprobs,
321                truncated: !tally.truncated_labels.is_empty(),
322                truncated_labels: tally.truncated_labels,
323                label_mass: tally.label_mass,
324            },
325
326            QuestionKind::Choice(spec) => {
327                let winner = math::argmax(&tally.probabilities);
328                Answer::Choice {
329                    choice: spec.options[winner].clone(),
330                    index: winner,
331                    confidence: tally.confidence,
332                    probabilities: spec
333                        .options
334                        .iter()
335                        .zip(&tally.probabilities)
336                        .map(|(option, p)| OptionProbability {
337                            option: option.clone(),
338                            probability: *p,
339                        })
340                        .collect(),
341                    raw_logprobs: tally.logprobs,
342                    truncated: !tally.truncated_labels.is_empty(),
343                    truncated_labels: tally.truncated_labels,
344                    label_mass: tally.label_mass,
345                }
346            }
347
348            QuestionKind::Score(spec) => {
349                let legend = spec.levels.legend();
350                let winner = math::argmax(&tally.probabilities);
351                Answer::Score {
352                    score: (winner + 1) as u8,
353                    expected_score: math::expected_level(&tally.probabilities),
354                    legend: legend[winner].clone(),
355                    confidence: tally.confidence,
356                    probabilities: legend
357                        .iter()
358                        .zip(&tally.probabilities)
359                        .enumerate()
360                        .map(|(i, (text, p))| LevelProbability {
361                            level: (i + 1) as u8,
362                            legend: text.clone(),
363                            probability: *p,
364                        })
365                        .collect(),
366                    raw_logprobs: tally.logprobs,
367                    truncated: !tally.truncated_labels.is_empty(),
368                    truncated_labels: tally.truncated_labels,
369                    label_mass: tally.label_mass,
370                }
371            }
372        };
373
374        Ok((answer, distribution))
375    }
376
377    /// The shared middle: prompt, one host call, read the labels back out.
378    async fn tally(
379        &self,
380        state: &str,
381        ballot: &Ballot<'_>,
382        model: &str,
383        calibration: Calibration,
384        question_id: &str,
385    ) -> Result<(Tally, FirstTokenDistribution), EngineError> {
386        let label_set = labels::labels(ballot.options.len());
387        let lettered: Vec<(&str, &str)> = label_set
388            .iter()
389            .copied()
390            .zip(ballot.options.iter().map(String::as_str))
391            .collect();
392
393        let distribution = self
394            .host
395            .first_token(FirstTokenRequest {
396                model: model.to_string(),
397                system: Some(prompt::SYSTEM.to_string()),
398                user: prompt::user_turn(state, ballot.question, &lettered),
399                top_logprobs: self.max_options(),
400                keep_alive: self.keep_alive.clone(),
401            })
402            .await
403            .map_err(|error| match error {
404                HostError::Status { status: 404, .. } => EngineError::UnknownModel {
405                    model: model.to_string(),
406                    host: self.host.name().to_string(),
407                    source: error,
408                },
409                other => EngineError::Host(other),
410            })?;
411
412        let tally = read_labels(&distribution, &label_set, calibration, question_id)?;
413        Ok((tally, distribution))
414    }
415}
416
417/// The first text that appears twice, compared the way the model sees it: trimmed.
418fn repeated(texts: &[String]) -> Option<&str> {
419    let mut seen = std::collections::BTreeSet::new();
420    texts.iter().map(|t| t.trim()).find(|t| !seen.insert(*t))
421}
422
423/// Reduce a question to its options. This is the only place the three primitives differ.
424fn ballot_for(kind: &QuestionKind) -> Ballot<'_> {
425    match kind {
426        // "Yes" first, so the noul probability is always option 0.
427        QuestionKind::Noul(question) => Ballot {
428            question: Some(question),
429            options: vec!["Yes".to_string(), "No".to_string()],
430        },
431        QuestionKind::Choice(spec) => Ballot {
432            question: asked(&spec.question),
433            options: spec.options.clone(),
434        },
435        QuestionKind::Score(spec) => Ballot {
436            question: asked(&spec.question),
437            options: spec.levels.legend(),
438        },
439    }
440}
441
442/// The question text, if there is one worth asking. A blank question is the same as none: it
443/// would otherwise put an empty `QUESTION:` line in front of the options.
444fn asked(question: &Option<String>) -> Option<&str> {
445    question.as_deref().filter(|q| !q.trim().is_empty())
446}
447
448/// Map a token distribution onto the labels we offered.
449fn read_labels(
450    distribution: &FirstTokenDistribution,
451    label_set: &[&str],
452    calibration: Calibration,
453    question_id: &str,
454) -> Result<Tally, EngineError> {
455    let mut logprobs = Vec::with_capacity(label_set.len());
456    let mut observed_any = false;
457    let mut truncated_labels = Vec::new();
458    // Observed labels only: a floor is a bound, and adding bounds would overstate the mass.
459    let mut label_mass = 0.0;
460
461    for label in label_set {
462        // Every token spelling this label counts toward it; see `math::logsumexp`.
463        let variants: Vec<f64> = distribution
464            .tokens
465            .iter()
466            .filter(|(token, _)| labels::matches(token, label))
467            .map(|(_, lp)| *lp)
468            .collect();
469
470        if variants.is_empty() {
471            // Ranked below everything the host reported, so the floor is a strict upper bound.
472            truncated_labels.push(label.to_string());
473            logprobs.push(distribution.floor);
474        } else {
475            observed_any = true;
476            let logprob = math::logsumexp(&variants);
477            label_mass += logprob.exp();
478            logprobs.push(logprob);
479        }
480    }
481
482    if !observed_any {
483        return Err(EngineError::NoLabelMatched {
484            question_id: question_id.to_string(),
485            expected: label_set.iter().map(|l| l.to_string()).collect(),
486            observed: distribution
487                .tokens
488                .iter()
489                .take(5)
490                .map(|(t, _)| t.clone())
491                .collect(),
492        });
493    }
494
495    let probabilities = math::softmax(&logprobs, calibration.temperature);
496    let confidence = math::confidence(&probabilities);
497
498    Ok(Tally {
499        logprobs: label_set
500            .iter()
501            .map(|l| l.to_string())
502            .zip(logprobs)
503            .collect(),
504        probabilities,
505        confidence,
506        truncated_labels,
507        // Reported logprobs over-sum a little from rounding; the share cannot pass 1.
508        label_mass: label_mass.min(1.0),
509    })
510}