1use 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 #[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 #[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
88struct Ballot<'a> {
90 question: Option<&'a str>,
91 options: Vec<String>,
93}
94
95struct Tally {
97 probabilities: Vec<f64>,
99 logprobs: BTreeMap<String, f64>,
102 confidence: f64,
103 truncated_labels: Vec<String>,
105 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 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 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 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 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 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 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
417fn 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
423fn ballot_for(kind: &QuestionKind) -> Ballot<'_> {
425 match kind {
426 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
442fn asked(question: &Option<String>) -> Option<&str> {
445 question.as_deref().filter(|q| !q.trim().is_empty())
446}
447
448fn 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 let mut label_mass = 0.0;
460
461 for label in label_set {
462 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 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 label_mass: label_mass.min(1.0),
509 })
510}