Skip to main content

wowlab_sentinel/http/
api_error.rs

1//! Typed HTTP error mapping for Sentinel routes.
2
3use 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}