Skip to main content

ntex/web/
state.rs

1use super::error::DefaultError;
2
3/// Application state used by web services.
4///
5/// A [`web::App`](super::App) is parameterized by a type that implements this
6/// trait. The state is available to handlers registered with
7/// [`Route::to_with_state()`](super::Route::to_with_state), and to filters,
8/// middleware, and services through their service context.
9///
10/// The associated `Error` type selects the application's error domain, which
11/// is used by [`WebError`](super::WebError) and
12/// [`WebResponseError`](super::WebResponseError). Most applications use
13/// [`DefaultError`].
14///
15/// The trait does not require `Clone`, but serving an application does:
16/// the server and [`HttpService`](crate::http::HttpService) clone the state
17/// for each connection.
18///
19/// ```rust
20/// use ntex::web;
21///
22/// #[derive(Clone)]
23/// struct ApplicationState {
24///     greeting: String,
25/// }
26///
27/// impl web::State for ApplicationState {
28///     type Error = web::DefaultError;
29/// }
30/// ```
31pub trait State: 'static {
32    /// Error domain used by the application.
33    type Error;
34}
35
36impl State for () {
37    type Error = DefaultError;
38}
39
40/// Wraps a value so it can be used as application state with [`DefaultError`].
41///
42/// `AppState<T>` implements [`State`] for any `'static` `T`, so a separate
43/// `State` implementation is not needed. It dereferences to the wrapped value
44/// and implements `Clone` and `Default` when `T` does.
45///
46/// ```rust
47/// use ntex::web;
48///
49/// #[derive(Clone)]
50/// struct Settings {
51///     service_name: &'static str,
52/// }
53///
54/// let state = web::AppState::new(Settings { service_name: "users" });
55/// assert_eq!(state.service_name, "users");
56/// ```
57#[derive(Clone, Default)]
58pub struct AppState<T> {
59    state: T,
60}
61
62impl<T> AppState<T> {
63    /// Creates application state that wraps `state`.
64    pub fn new(state: T) -> Self {
65        AppState { state }
66    }
67
68    /// Returns a reference to the wrapped value.
69    pub fn st(&self) -> &T {
70        &self.state
71    }
72}
73
74impl<T: 'static> State for AppState<T> {
75    type Error = DefaultError;
76}
77
78impl<T> std::ops::Deref for AppState<T> {
79    type Target = T;
80
81    fn deref(&self) -> &T {
82        &self.state
83    }
84}
85
86#[cfg(test)]
87mod tests {
88    use std::{cell::Cell, convert::Infallible, rc::Rc};
89
90    use super::*;
91    use crate::http::{Method, StatusCode};
92    use crate::service::ServiceFactory;
93    use crate::web::test::{TestRequest, init_service_st, read_body};
94    use crate::web::{self, App, HttpRequest, WebRequest, types::Path};
95
96    #[derive(Clone, Default)]
97    struct Counter {
98        name: &'static str,
99        hits: Rc<Cell<usize>>,
100    }
101
102    type St = AppState<Counter>;
103
104    async fn index(st: &St, (): ()) -> String {
105        st.hits.set(st.hits.get() + 1);
106        format!("index:{}", st.name)
107    }
108
109    async fn item(st: &St, (): (), id: Path<u32>) -> String {
110        format!("item:{}:{}", st.st().name, id.into_inner())
111    }
112
113    async fn fallback(st: &St, (): (), req: HttpRequest) -> String {
114        format!("fallback:{}:{}", st.name, req.method())
115    }
116
117    async fn filtered(_: &St, n: usize) -> String {
118        format!("filtered:{n}")
119    }
120
121    #[test]
122    fn app_state() {
123        let st = AppState::new(Counter {
124            name: "test",
125            ..Counter::default()
126        });
127        assert_eq!(st.st().name, "test");
128        assert_eq!(st.name, "test");
129        assert_eq!(st.clone().name, "test");
130        assert_eq!(AppState::<Counter>::default().name, "");
131    }
132
133    #[crate::rt_test]
134    async fn state_handlers() {
135        let st = AppState::new(Counter {
136            name: "app",
137            ..Counter::default()
138        });
139        let srv = init_service_st(
140            st.clone(),
141            App::<St>::new()
142                .service(web::resource("/index").to_with_state(index))
143                .service(web::resource("/item/{id}").to_with_state(item))
144                .service(
145                    web::resource("/multi")
146                        .route(web::get().to(async || "get"))
147                        .route(web::delete().to_with_state(index))
148                        .to(async || "any"),
149                )
150                .service(
151                    web::resource("/fallback")
152                        .route(web::get().to(async || "get"))
153                        .to_with_state(fallback),
154                )
155                .service(
156                    web::resource("/filtered")
157                        .filter(async |req: WebRequest<()>| {
158                            Ok::<_, Infallible>(req.map_state(|()| 7usize))
159                        })
160                        .to_with_state(filtered),
161                )
162                .route("/route", web::to_with_state(index))
163                .build(),
164        )
165        .await;
166
167        let check = async |method: Method, uri: &str, status: StatusCode, body: &str| {
168            let req = TestRequest::with_uri(uri).method(method).to_request();
169            let resp = srv.call(req).await.unwrap();
170            assert_eq!(resp.status(), status, "{uri}");
171            assert_eq!(read_body(resp).await, body.as_bytes(), "{uri}");
172        };
173
174        check(Method::GET, "/index", StatusCode::OK, "index:app").await;
175        check(Method::GET, "/item/10", StatusCode::OK, "item:app:10").await;
176        check(
177            Method::GET,
178            "/item/abc",
179            StatusCode::NOT_FOUND,
180            "Path deserialize error: can not parse \"abc\" to a u32",
181        )
182        .await;
183        check(Method::GET, "/multi", StatusCode::OK, "get").await;
184        check(Method::DELETE, "/multi", StatusCode::OK, "index:app").await;
185        check(Method::POST, "/multi", StatusCode::OK, "any").await;
186        check(Method::GET, "/fallback", StatusCode::OK, "get").await;
187        check(
188            Method::POST,
189            "/fallback",
190            StatusCode::OK,
191            "fallback:app:POST",
192        )
193        .await;
194        check(Method::GET, "/filtered", StatusCode::OK, "filtered:7").await;
195        check(Method::PUT, "/route", StatusCode::OK, "index:app").await;
196        assert_eq!(st.hits.get(), 3);
197    }
198
199    #[crate::rt_test]
200    async fn build_with() {
201        let st = AppState::new(Counter {
202            name: "fixed",
203            ..Counter::default()
204        });
205        let srv = App::<St>::new()
206            .route("/", web::get().to_with_state(index))
207            .build_with::<usize>(st.clone())
208            .pipeline(1usize)
209            .await
210            .unwrap();
211
212        let resp = srv.call(TestRequest::default().to_request()).await.unwrap();
213        assert_eq!(resp.status(), StatusCode::OK);
214        let resp = srv
215            .call(TestRequest::with_uri("/missing").to_request())
216            .await
217            .unwrap();
218        assert_eq!(resp.status(), StatusCode::NOT_FOUND);
219        assert_eq!(st.hits.get(), 1);
220    }
221}