1use super::error::DefaultError;
2
3pub trait State: 'static {
32 type Error;
34}
35
36impl State for () {
37 type Error = DefaultError;
38}
39
40#[derive(Clone, Default)]
58pub struct AppState<T> {
59 state: T,
60}
61
62impl<T> AppState<T> {
63 pub fn new(state: T) -> Self {
65 AppState { state }
66 }
67
68 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}