1use axum::Json;
4use axum::extract::rejection::JsonRejection;
5use axum::extract::{FromRequest, Request};
6use axum::http::StatusCode;
7use axum::response::{IntoResponse, Response};
8use cerno_core::EngineError;
9use cerno_types::{ErrorCode, ErrorResponse};
10
11pub fn status_for(code: ErrorCode) -> StatusCode {
19 match code {
20 ErrorCode::InvalidRequest => StatusCode::BAD_REQUEST,
22
23 ErrorCode::TooManyOptions
24 | ErrorCode::TooManyQuestions
25 | ErrorCode::TooFewOptions
26 | ErrorCode::InvalidLevels
27 | ErrorCode::EmptyState
28 | ErrorCode::EmptyQuestion
29 | ErrorCode::DuplicateQuestionId
30 | ErrorCode::EmptyQuestionId
31 | ErrorCode::DuplicateOption
32 | ErrorCode::NoQuestions
33 | ErrorCode::UnknownModel
34 | ErrorCode::InvalidCalibration => StatusCode::UNPROCESSABLE_ENTITY,
35
36 ErrorCode::NoLabelMatched | ErrorCode::HostUnavailable => StatusCode::BAD_GATEWAY,
37 ErrorCode::HostTimeout => StatusCode::GATEWAY_TIMEOUT,
38 ErrorCode::Internal => StatusCode::INTERNAL_SERVER_ERROR,
39 }
40}
41
42pub struct ApiJson<T>(pub T);
48
49impl<T, S> FromRequest<S> for ApiJson<T>
50where
51 Json<T>: FromRequest<S, Rejection = JsonRejection>,
52 S: Send + Sync,
53{
54 type Rejection = ApiError;
55
56 async fn from_request(req: Request, state: &S) -> Result<Self, Self::Rejection> {
57 match Json::<T>::from_request(req, state).await {
58 Ok(Json(value)) => Ok(Self(value)),
59 Err(rejection) => Err(ApiError::new(
60 ErrorCode::InvalidRequest,
61 rejection.body_text(),
62 )),
63 }
64 }
65}
66
67pub struct ApiError {
69 pub code: ErrorCode,
70 pub message: String,
71 pub question_id: Option<String>,
72}
73
74impl ApiError {
75 pub fn new(code: ErrorCode, message: impl Into<String>) -> Self {
76 Self {
77 code,
78 message: message.into(),
79 question_id: None,
80 }
81 }
82}
83
84impl From<EngineError> for ApiError {
85 fn from(err: EngineError) -> Self {
86 Self {
87 code: err.code(),
88 question_id: err.question_id().map(str::to_string),
89 message: err.to_string(),
90 }
91 }
92}
93
94impl IntoResponse for ApiError {
95 fn into_response(self) -> Response {
96 let status = status_for(self.code);
97
98 if status.is_server_error() {
99 tracing::error!(code = ?self.code, question_id = ?self.question_id, "{}", self.message);
100 } else {
101 tracing::debug!(code = ?self.code, question_id = ?self.question_id, "{}", self.message);
102 }
103
104 (
105 status,
106 Json(ErrorResponse {
107 code: self.code,
108 message: self.message,
109 question_id: self.question_id,
110 }),
111 )
112 .into_response()
113 }
114}
115
116#[cfg(test)]
117mod tests {
118 use super::*;
119
120 #[test]
121 fn caller_mistakes_are_unprocessable() {
122 for code in [
123 ErrorCode::TooManyOptions,
124 ErrorCode::TooManyQuestions,
125 ErrorCode::TooFewOptions,
126 ErrorCode::InvalidLevels,
127 ErrorCode::EmptyState,
128 ErrorCode::NoQuestions,
129 ErrorCode::DuplicateQuestionId,
130 ErrorCode::EmptyQuestionId,
131 ErrorCode::DuplicateOption,
132 ErrorCode::InvalidCalibration,
133 ErrorCode::UnknownModel,
134 ] {
135 assert_eq!(
136 status_for(code),
137 StatusCode::UNPROCESSABLE_ENTITY,
138 "{code:?}"
139 );
140 }
141 }
142
143 #[test]
145 fn model_and_host_failures_are_gateway_errors() {
146 assert_eq!(
147 status_for(ErrorCode::NoLabelMatched),
148 StatusCode::BAD_GATEWAY
149 );
150 assert_eq!(
151 status_for(ErrorCode::HostUnavailable),
152 StatusCode::BAD_GATEWAY
153 );
154 assert_eq!(
155 status_for(ErrorCode::HostTimeout),
156 StatusCode::GATEWAY_TIMEOUT
157 );
158 }
159
160 #[test]
161 fn a_body_that_is_not_a_request_is_a_bad_request() {
162 assert_eq!(
163 status_for(ErrorCode::InvalidRequest),
164 StatusCode::BAD_REQUEST
165 );
166 }
167
168 #[test]
169 fn a_failure_inside_cerno_is_an_internal_error() {
170 assert_eq!(
171 status_for(ErrorCode::Internal),
172 StatusCode::INTERNAL_SERVER_ERROR
173 );
174 }
175}