Skip to main content

toolkit/api/rest/extract/
query.rs

1//! `Query<T>` - a drop-in `axum::extract::Query<T>` replacement whose
2//! extraction failures render as `application/problem+json` (RFC 9457
3//! `Problem`) instead of axum's default plain-text rejection body. Same
4//! class of gap as `extract::Json`, same fix shape - see this module's
5//! parent doc comment.
6
7use axum::extract::rejection::QueryRejection;
8use axum::extract::{FromRequestParts, Query as AxumQuery};
9use axum::http::request::Parts;
10use toolkit_canonical_errors::CanonicalError;
11
12use super::error::rejection_to_canonical;
13
14/// Drop-in replacement for `axum::extract::Query<T>` as a handler parameter.
15/// Extraction success is identical to `axum::extract::Query<T>`; extraction
16/// failure produces a `CanonicalError` instead of `QueryRejection`'s
17/// plain-text body.
18#[derive(Debug, Clone, Copy, Default)]
19pub struct Query<T>(pub T);
20
21impl<T, S> FromRequestParts<S> for Query<T>
22where
23    AxumQuery<T>: FromRequestParts<S, Rejection = QueryRejection>,
24    S: Send + Sync,
25{
26    type Rejection = CanonicalError;
27
28    async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
29        match AxumQuery::<T>::from_request_parts(parts, state).await {
30            Ok(AxumQuery(value)) => Ok(Self(value)),
31            Err(rejection) => Err(query_rejection_to_canonical(&rejection)),
32        }
33    }
34}
35
36/// Maps a `QueryRejection` to a `CanonicalError`. Not a `From` impl - both
37/// types are foreign to this crate, so the orphan rule forbids it.
38///
39/// `QueryRejection` is `#[non_exhaustive]` (today it has exactly one
40/// variant, `FailedToDeserializeQueryString`), but unlike
41/// `extract::json`, the mapping below was never variant-specific in the
42/// first place - it's driven entirely by `.status()`/`.body_text()` with one
43/// fixed code. An unknown future variant is handled correctly by that same
44/// logic with nothing invented, so there's no panic - only a log line for
45/// visibility when the variant isn't the one this function was written
46/// against.
47fn query_rejection_to_canonical(rejection: &QueryRejection) -> CanonicalError {
48    if !matches!(rejection, QueryRejection::FailedToDeserializeQueryString(_)) {
49        tracing::error!(
50            rejection = %rejection,
51            "extract::Query: unhandled QueryRejection variant, falling back to the same status-driven classification"
52        );
53    }
54    rejection_to_canonical(
55        "query",
56        "invalid_query_string",
57        rejection.status().as_u16(),
58        rejection.body_text(),
59    )
60}
61
62#[cfg(test)]
63#[cfg_attr(coverage_nightly, coverage(off))]
64mod tests {
65    use axum::Router;
66    use axum::body::Body;
67    use axum::http::{Request, StatusCode, header};
68    use axum::routing::get;
69    use serde::Deserialize;
70    use serde_json::{Value, json};
71    use tower::ServiceExt;
72
73    use super::Query;
74
75    const RESOURCE_TYPE: &str = "gts.cf.core.http.request.v1~";
76    const INVALID_ARGUMENT_TYPE: &str =
77        "gts://gts.cf.core.errors.err.v1~cf.core.err.invalid_argument.v1~";
78
79    #[derive(Debug, Deserialize)]
80    struct Filter {
81        page: u32,
82    }
83
84    fn app() -> Router {
85        Router::new().route(
86            "/items",
87            get(|Query(f): Query<Filter>| async move { axum::Json(f.page) }),
88        )
89    }
90
91    async fn get_uri(uri: &str) -> axum::response::Response {
92        let req = Request::builder().uri(uri).body(Body::empty()).unwrap();
93        app().oneshot(req).await.unwrap()
94    }
95
96    async fn body_json(response: axum::response::Response) -> Value {
97        let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
98            .await
99            .unwrap();
100        serde_json::from_slice(&bytes).expect("response body is valid JSON")
101    }
102
103    #[tokio::test]
104    async fn valid_query_extracts_normally() {
105        // Echoes the extracted `page` back and asserts it, not just the
106        // status - a `Query<T>` that silently yielded a wrong or defaulted
107        // value would still return 200 and pass a status-only assertion.
108        let res = get_uri("/items?page=1").await;
109        assert_eq!(res.status(), StatusCode::OK);
110        let json = body_json(res).await;
111        assert_eq!(json, json!(1));
112    }
113
114    #[tokio::test]
115    async fn non_numeric_field_returns_400_problem() {
116        let res = get_uri("/items?page=not-a-number").await;
117        assert_eq!(res.status(), StatusCode::BAD_REQUEST);
118        assert_eq!(
119            res.headers().get(header::CONTENT_TYPE).unwrap(),
120            "application/problem+json"
121        );
122        let json = body_json(res).await;
123        assert_eq!(
124            json,
125            json!({
126                "type": INVALID_ARGUMENT_TYPE,
127                "title": "Invalid Argument",
128                "status": 400,
129                "detail": "Request validation failed",
130                "context": {
131                    "resource_type": RESOURCE_TYPE,
132                    "field_violations": [{
133                        "field": "query",
134                        "description": "Failed to deserialize query string: page: invalid digit found in string",
135                        "reason": "invalid_query_string",
136                    }],
137                },
138            })
139        );
140    }
141
142    #[tokio::test]
143    async fn missing_required_field_returns_400_problem() {
144        let res = get_uri("/items").await;
145        assert_eq!(res.status(), StatusCode::BAD_REQUEST);
146        assert_eq!(
147            res.headers().get(header::CONTENT_TYPE).unwrap(),
148            "application/problem+json"
149        );
150        let json = body_json(res).await;
151        assert_eq!(
152            json,
153            json!({
154                "type": INVALID_ARGUMENT_TYPE,
155                "title": "Invalid Argument",
156                "status": 400,
157                "detail": "Request validation failed",
158                "context": {
159                    "resource_type": RESOURCE_TYPE,
160                    "field_violations": [{
161                        "field": "query",
162                        "description": "Failed to deserialize query string: missing field `page`",
163                        "reason": "invalid_query_string",
164                    }],
165                },
166            })
167        );
168    }
169}