Skip to main content

cerno_server/
routes.rs

1//! HTTP handlers.
2
3use crate::error::{ApiError, ApiJson};
4use crate::state::AppState;
5use axum::Json;
6use axum::extract::State;
7use cerno_types::{
8    Answer, Calibration, ErrorCode, ErrorResponse, ModelInfo, ModelsResponse, SystemOneRequest,
9    SystemOneResponse, Timing, Usage,
10};
11use std::collections::{BTreeMap, HashMap};
12use std::sync::Arc;
13use std::time::Instant;
14use tokio::task::JoinSet;
15
16/// Liveness. Registered outside the traced router, so health checks never reach the logs.
17pub async fn health() -> &'static str {
18    "ok"
19}
20
21/// The models this server will answer for.
22#[utoipa::path(
23    get,
24    path = "/v1/models",
25    responses((status = 200, body = ModelsResponse)),
26    tag = "cerno",
27)]
28pub async fn models(State(state): State<AppState>) -> Json<ModelsResponse> {
29    Json(ModelsResponse {
30        models: state
31            .config
32            .models
33            .iter()
34            .map(|(alias, entry)| ModelInfo {
35                alias: alias.clone(),
36                model: entry.model.clone(),
37                calibration: entry.calibration.into(),
38            })
39            .collect(),
40        default: state.config.default_model.clone(),
41    })
42}
43
44/// Ask several questions about one state.
45///
46/// Every question is one forward pass, and they are independent, so they run concurrently up to
47/// the configured limit rather than one after another.
48///
49/// The request succeeds or fails as a whole: if any question fails, the error names it and the
50/// answers to the others are discarded. The same holds when the whole request, waiting for a
51/// free slot included, outlasts the server's request timeout: that is a 504 `host_timeout`.
52#[utoipa::path(
53    post,
54    path = "/v1/systemone",
55    request_body = SystemOneRequest,
56    responses(
57        (status = 200, body = SystemOneResponse),
58        (status = 400, body = ErrorResponse, description = "the body is not a valid request"),
59        (status = 422, body = ErrorResponse, description = "the request cannot be answered as written"),
60        (status = 502, body = ErrorResponse, description = "the model or its runtime failed"),
61        (status = 504, body = ErrorResponse, description = "the host timed out"),
62        (status = 500, body = ErrorResponse, description = "a failure inside cerno"),
63    ),
64    tag = "cerno",
65)]
66pub async fn systemone(
67    State(state): State<AppState>,
68    ApiJson(request): ApiJson<SystemOneRequest>,
69) -> Result<Json<SystemOneResponse>, ApiError> {
70    let started = Instant::now();
71
72    let requested = request
73        .model
74        .clone()
75        .unwrap_or_else(|| state.config.default_model.clone());
76
77    let (model, model_calibration) = state.config.resolve(&requested).ok_or_else(|| {
78        ApiError::new(
79            ErrorCode::UnknownModel,
80            format!(
81                "model {requested:?} is not configured and strict_models is on; known: {:?}",
82                state.config.known_models()
83            ),
84        )
85    })?;
86
87    // Request beats alias beats identity.
88    let calibration: Calibration = request.calibration.unwrap_or(model_calibration);
89
90    // Validate every question before loading a model, so a bad request costs nothing.
91    state.engine.validate(&request)?;
92
93    // Shared by every question's task rather than copied into each: a state can be large.
94    let state_text: Arc<str> = Arc::from(request.state.as_str());
95    let deadline = started + state.config.request_timeout;
96
97    let mut tasks = JoinSet::new();
98    // Which question each task answers. A task that panics comes back as a bare JoinError, so
99    // the id has to be recoverable from the task alone.
100    let mut question_of = HashMap::new();
101    for question in request.questions.clone() {
102        let engine = state.engine.clone();
103        let semaphore = state.semaphore.clone();
104        let state_text = state_text.clone();
105        let model = model.clone();
106
107        let id = question.id.clone();
108        let handle = tasks.spawn(async move {
109            let _permit = semaphore
110                .acquire_owned()
111                .await
112                .expect("semaphore is never closed");
113            engine
114                .answer(&state_text, &question, &model, calibration)
115                .await
116        });
117        question_of.insert(handle.id(), id);
118    }
119
120    let mut answers: BTreeMap<String, Answer> = BTreeMap::new();
121    let mut input_tokens: u32 = 0;
122
123    // Returning early drops the JoinSet, which aborts every question still running or queued,
124    // and gives their permits back.
125    let out_of_time = || ApiError {
126        code: ErrorCode::HostTimeout,
127        message: format!(
128            "the request did not finish within {:?}; ask fewer questions at once, or raise \
129             CERNO_REQUEST_TIMEOUT_SECS",
130            state.config.request_timeout
131        ),
132        question_id: None,
133    };
134
135    while let Some(joined) = tokio::time::timeout_at(deadline.into(), tasks.join_next_with_id())
136        .await
137        .map_err(|_| out_of_time())?
138    {
139        let (task, result) = joined.map_err(|e| {
140            let id = question_of.get(&e.id()).cloned();
141            ApiError {
142                code: ErrorCode::Internal,
143                message: format!("answering question {id:?} failed to complete: {e}"),
144                question_id: id,
145            }
146        })?;
147
148        let (answer, tokens) = result?;
149        let id = question_of
150            .remove(&task)
151            .expect("every task was registered when it was spawned");
152        input_tokens = input_tokens.saturating_add(tokens);
153        answers.insert(id, answer);
154    }
155
156    Ok(Json(SystemOneResponse {
157        usage: Usage {
158            input_tokens,
159            questions: answers.len(),
160        },
161        answers,
162        model,
163        timing_ms: Timing {
164            total: started.elapsed().as_millis() as u64,
165        },
166    }))
167}