1use axum::{
4 Json,
5 http::StatusCode,
6 response::{IntoResponse, Response},
7};
8use serde::Serialize;
9
10use super::{
11 auth::AuthError,
12 services::{
13 chunks::{CompletionError, CompletionPersistenceError},
14 nodes::NodeOperationError,
15 work_context::WorkContextError,
16 },
17};
18
19#[derive(Debug, thiserror::Error)]
20pub(super) enum ApiError {
21 #[error("{0}")]
22 Auth(
23 #[from]
24 #[source]
25 AuthError,
26 ),
27 #[error("{0}")]
28 Node(
29 #[from]
30 #[source]
31 NodeOperationError,
32 ),
33 #[error("{0}")]
34 WorkContext(
35 #[from]
36 #[source]
37 WorkContextError,
38 ),
39 #[error("{0}")]
40 Completion(
41 #[from]
42 #[source]
43 CompletionError,
44 ),
45 #[error("invalid protobuf payload")]
46 InvalidProtobuf(#[source] prost::DecodeError),
47 #[error("invalid webhook payload")]
48 InvalidWebhookPayload(#[source] serde_json::Error),
49}
50
51#[derive(Debug, Serialize)]
52struct ApiErrorBody {
53 error: String,
54 #[serde(skip_serializing_if = "Option::is_none")]
55 details: Option<String>,
56}
57
58impl ApiError {
59 pub(super) const fn invalid_protobuf(source: prost::DecodeError) -> Self {
60 Self::InvalidProtobuf(source)
61 }
62
63 pub(super) const fn invalid_webhook_payload(source: serde_json::Error) -> Self {
64 Self::InvalidWebhookPayload(source)
65 }
66
67 fn status_and_body(&self) -> (StatusCode, ApiErrorBody) {
68 match self {
69 Self::Auth(error) => error_body(StatusCode::UNAUTHORIZED, error.wire_message()),
70 Self::Node(error) => node_error_body(error),
71 Self::WorkContext(error) => work_context_error_body(error),
72 Self::Completion(error) => completion_error_body(error),
73 Self::InvalidProtobuf(source) => (
74 StatusCode::BAD_REQUEST,
75 ApiErrorBody {
76 error: "Invalid protobuf payload".to_string(),
77 details: Some(source.to_string()),
78 },
79 ),
80 Self::InvalidWebhookPayload(_) => {
81 error_body(StatusCode::BAD_REQUEST, "Invalid payload")
82 }
83 }
84 }
85}
86
87fn node_error_body(error: &NodeOperationError) -> (StatusCode, ApiErrorBody) {
88 match error {
89 NodeOperationError::InvalidToken { .. } => {
90 error_body(StatusCode::UNAUTHORIZED, "Invalid token")
91 }
92 NodeOperationError::CreateNode { .. } => {
93 error_body(StatusCode::BAD_REQUEST, "Failed to create node")
94 }
95 NodeOperationError::TokenGeneration(_) => {
96 error_body(StatusCode::INTERNAL_SERVER_ERROR, "Token generation failed")
97 }
98 NodeOperationError::NodeNotClaimed { .. } => {
99 error_body(StatusCode::NOT_FOUND, "Node not found or not claimed")
100 }
101 NodeOperationError::NodeNotFound => error_body(StatusCode::NOT_FOUND, "Node not found"),
102 NodeOperationError::Database(_) => {
103 error_body(StatusCode::INTERNAL_SERVER_ERROR, "Database error")
104 }
105 }
106}
107
108fn work_context_error_body(error: &WorkContextError) -> (StatusCode, ApiErrorBody) {
109 match error {
110 WorkContextError::InvalidHash(_) => {
111 error_body(StatusCode::BAD_REQUEST, "Invalid work context hash")
112 }
113 WorkContextError::NoMatchingClaim => error_body(StatusCode::FORBIDDEN, "No matching claim"),
114 }
115}
116
117fn completion_error_body(error: &CompletionError) -> (StatusCode, ApiErrorBody) {
118 match error {
119 CompletionError::InvalidJobId(_) => error_body(StatusCode::BAD_REQUEST, "Invalid job_id"),
120 CompletionError::InvalidWorkContextHashLength { .. } => {
121 error_body(StatusCode::BAD_REQUEST, "Invalid work_context_hash length")
122 }
123 CompletionError::JobNotFound => {
124 error_body(StatusCode::NOT_FOUND, "Job not found or not running")
125 }
126 CompletionError::Ingest(source) => {
127 error_body(StatusCode::INTERNAL_SERVER_ERROR, source.to_string())
128 }
129 CompletionError::Conflict(message) => error_body(StatusCode::CONFLICT, message),
130 CompletionError::Stale(message) => error_body(StatusCode::GONE, message),
131 CompletionError::Finalize(_) => {
132 error_body(StatusCode::INTERNAL_SERVER_ERROR, "Finalize failed")
133 }
134 CompletionError::Persistence(error) => persistence_error_body(error),
135 }
136}
137
138fn persistence_error_body(error: &CompletionPersistenceError) -> (StatusCode, ApiErrorBody) {
139 let message = match error {
140 CompletionPersistenceError::Begin(_) => "Database error",
141 CompletionPersistenceError::Save(_) => "Failed to save result",
142 CompletionPersistenceError::Commit(_) => "Failed to commit",
143 };
144
145 error_body(StatusCode::INTERNAL_SERVER_ERROR, message)
146}
147
148fn error_body(status: StatusCode, message: impl Into<String>) -> (StatusCode, ApiErrorBody) {
149 (
150 status,
151 ApiErrorBody {
152 error: message.into(),
153 details: None,
154 },
155 )
156}
157
158impl IntoResponse for ApiError {
159 fn into_response(self) -> Response {
160 let (status, body) = self.status_and_body();
161
162 (status, Json(body)).into_response()
163 }
164}
165
166#[cfg(test)]
167mod tests {
168 use std::error::Error;
169
170 use googletest::prelude::*;
171
172 use super::*;
173
174 async fn body(error: ApiError) -> (StatusCode, String) {
175 let response = error.into_response();
176 let status = response.status();
177 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
178 .await
179 .expect("error body");
180
181 (
182 status,
183 String::from_utf8(bytes.to_vec()).expect("UTF-8 JSON"),
184 )
185 }
186
187 #[gtest]
188 #[tokio::test]
189 async fn service_errors_preserve_exact_status_and_wire_body() -> Result<()> {
190 let cases = [
191 (
192 ApiError::from(NodeOperationError::InvalidToken { source: None }),
193 StatusCode::UNAUTHORIZED,
194 r#"{"error":"Invalid token"}"#,
195 ),
196 (
197 ApiError::from(WorkContextError::NoMatchingClaim),
198 StatusCode::FORBIDDEN,
199 r#"{"error":"No matching claim"}"#,
200 ),
201 (
202 ApiError::from(CompletionError::Conflict("claim mismatch".into())),
203 StatusCode::CONFLICT,
204 r#"{"error":"claim mismatch"}"#,
205 ),
206 (
207 ApiError::from(CompletionError::Stale("claim stale".into())),
208 StatusCode::GONE,
209 r#"{"error":"claim stale"}"#,
210 ),
211 ];
212
213 for (error, expected_status, expected_body) in cases {
214 let (status, actual_body) = body(error).await;
215
216 verify_eq!(status, expected_status)?;
217 verify_eq!(actual_body, expected_body)?;
218 }
219
220 Ok(())
221 }
222
223 #[gtest]
224 #[tokio::test]
225 async fn protobuf_error_preserves_details_and_source_chain() -> Result<()> {
226 let source =
227 <wowlab_types::proto::BatchChunkCompletion as prost::Message>::decode(&[0xff][..])
228 .err()
229 .or_fail()?;
230 let error = ApiError::invalid_protobuf(source);
231
232 verify_true!(error.source().is_some())?;
233
234 let (status, actual_body) = body(error).await;
235
236 verify_eq!(status, StatusCode::BAD_REQUEST)?;
237 verify_true!(
238 actual_body.starts_with(r#"{"error":"Invalid protobuf payload","details":""#)
239 )?;
240
241 Ok(())
242 }
243
244 #[gtest]
245 #[tokio::test]
246 async fn completion_ingest_error_preserves_source_chain_and_wire_message() -> Result<()> {
247 let source = <wowlab_types::proto::ChunkTelemetry as prost::Message>::decode(&[0xff][..])
248 .err()
249 .or_fail()?;
250 let error = ApiError::from(CompletionError::Ingest(
251 crate::strategy::StrategyError::from(source),
252 ));
253 let completion = error.source().or_fail()?;
254 let strategy = completion.source().or_fail()?;
255
256 verify_true!(
257 strategy
258 .source()
259 .and_then(Error::source)
260 .is_some_and(<dyn Error>::is::<prost::DecodeError>)
261 )?;
262
263 let (status, actual_body) = body(error).await;
264
265 verify_eq!(status, StatusCode::INTERNAL_SERVER_ERROR)?;
266 verify_eq!(
267 actual_body,
268 r#"{"error":"state decode error: failed to decode Protobuf message: invalid varint"}"#
269 )?;
270
271 Ok(())
272 }
273
274 #[gtest]
275 #[tokio::test]
276 async fn database_error_keeps_source_while_hiding_it_from_wire_body() -> Result<()> {
277 let error = ApiError::from(NodeOperationError::Database(sqlx::Error::Protocol(
278 "database-secret".into(),
279 )));
280
281 verify_eq!(
282 error.source().map(ToString::to_string),
283 Some("database error".to_string())
284 )?;
285 verify_true!(
286 error
287 .source()
288 .and_then(Error::source)
289 .is_some_and(|source| source.to_string().contains("database-secret"))
290 )?;
291
292 let (status, actual_body) = body(error).await;
293
294 verify_eq!(status, StatusCode::INTERNAL_SERVER_ERROR)?;
295 verify_eq!(actual_body, r#"{"error":"Database error"}"#)?;
296
297 Ok(())
298 }
299}