1use 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
16pub async fn health() -> &'static str {
18 "ok"
19}
20
21#[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#[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 let calibration: Calibration = request.calibration.unwrap_or(model_calibration);
89
90 state.engine.validate(&request)?;
92
93 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 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 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}