1use std::{cell::Cell, fmt, marker};
2
3use crate::ctx::{Ctx, WaitersRef};
4use crate::{IntoService, IntoServiceFactory, Service, ServiceFactory};
5use crate::{ServiceCaller, ServiceChain, ServiceChainFactory};
6
7pub fn apply_fn<S, St, Req, F, In, Out, Err>(
12 service: impl IntoService<S, St, Req>,
13 f: F,
14) -> ServiceChain<Apply<S, St, Req, F, In, Out, Err>, St, In>
15where
16 S: Service<St, Req>,
17 F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err>,
18 Err: From<S::Error>,
19{
20 crate::service(Apply::new(service.into_service(), f))
21}
22
23pub fn apply_fn_factory<Sf, St, Req, F, In, Out, Err>(
25 service: impl IntoServiceFactory<Sf, St, Req>,
26 f: F,
27) -> ServiceChainFactory<ApplyFactory<F, Sf, St, Req, In, Out, Err>, St, In>
28where
29 Sf: ServiceFactory<St, Req>,
30 F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
31 Err: From<Sf::Error>,
32{
33 crate::factory(ApplyFactory::new(service.into_factory(), f))
34}
35
36#[derive(Debug)]
37pub struct ApplyCtx<'a, S, St, Req> {
39 idx: u32,
40 waiters: &'a WaitersRef,
41 service: &'a S,
42 st: &'a St,
43 ready: &'a Cell<bool>,
44 entered: &'a Cell<u32>,
45 r: marker::PhantomData<Req>,
46}
47
48impl<S: Service<St, Req>, St, Req> ApplyCtx<'_, S, St, Req> {
49 #[inline]
51 pub fn st(&self) -> &St {
52 self.st
53 }
54
55 #[inline]
57 pub async fn call(&self, req: Req) -> Result<S::Res, S::Error> {
58 self.enter(req, self.ready.get()).await
59 }
60
61 async fn enter(&self, req: Req, skip_ready: bool) -> Result<S::Res, S::Error> {
62 let ctx = Ctx::<S, St>::new(self.idx, self.waiters, self.st);
63 if !skip_ready {
64 ctx.ready(self.service).await?;
65 }
66 self.ready.set(false);
68 self.entered.set(self.entered.get().wrapping_add(1));
69 ctx.call_nowait(self.service, req).await
70 }
71}
72
73impl<S: Service<St, Req>, St, Req> ServiceCaller<Req, S::Res, S::Error>
74 for ApplyCtx<'_, S, St, Req>
75{
76 #[inline]
77 async fn call_service(&self, req: Req) -> Result<S::Res, S::Error> {
78 self.enter(req, false).await
79 }
80}
81
82pub struct Apply<S, St, Req, F, In, Out, Err> {
84 svc: S,
85 f: F,
86 ready: Cell<bool>,
87 entered: Cell<u32>,
88 r: marker::PhantomData<fn(St, Req) -> (In, Out, Err)>,
89}
90
91impl<S, St, Req, F, In, Out, Err> Apply<S, St, Req, F, In, Out, Err>
92where
93 F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err>,
94{
95 pub(crate) fn new(svc: S, f: F) -> Self {
96 Apply {
97 f,
98 svc,
99 ready: Cell::new(false),
100 entered: Cell::new(0),
101 r: marker::PhantomData,
102 }
103 }
104}
105
106impl<S, St, Req, F, In, Out, Err> Clone for Apply<S, St, Req, F, In, Out, Err>
107where
108 S: Clone,
109 F: Clone,
110{
111 fn clone(&self) -> Self {
112 Apply {
113 svc: self.svc.clone(),
114 f: self.f.clone(),
115 ready: Cell::new(false),
116 entered: Cell::new(0),
117 r: marker::PhantomData,
118 }
119 }
120}
121
122impl<S, St, Req, F, In, Out, Err> fmt::Debug for Apply<S, St, Req, F, In, Out, Err>
123where
124 S: fmt::Debug,
125{
126 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
127 f.debug_struct("Apply")
128 .field("svc", &self.svc)
129 .field("map", &std::any::type_name::<F>())
130 .finish()
131 }
132}
133
134impl<S, St, Req, F, In, Out, Err> Service<St, In> for Apply<S, St, Req, F, In, Out, Err>
135where
136 S: Service<St, Req>,
137 F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err>,
138 Err: From<S::Error>,
139{
140 type Res = Out;
141 type Error = Err;
142
143 #[inline]
144 async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), Err> {
145 let entered = self.entered.get();
146 let result = ctx.ready(&self.svc).await.map_err(From::from);
147 if result.is_err() {
148 self.ready.set(false);
149 } else if self.entered.get() == entered {
150 self.ready.set(true);
152 }
153 result
154 }
155
156 #[inline]
157 async fn call(&self, req: In, ctx: Ctx<'_, Self, St>) -> Result<Out, Err> {
158 let (idx, waiters, st) = ctx.inner();
159
160 let ctx = ApplyCtx {
161 idx,
162 waiters,
163 st,
164 ready: &self.ready,
165 entered: &self.entered,
166 service: &self.svc,
167 r: marker::PhantomData,
168 };
169 (self.f)(req, &ctx).await
170 }
171
172 crate::forward_shutdown!(St, svc);
173}
174
175pub struct ApplyFactory<F, Sf, St, Req, In, Out, Err>
177where
178 F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
179 Sf: ServiceFactory<St, Req>,
180{
181 f: F,
182 sf: Sf,
183 r: marker::PhantomData<fn(St, Req) -> (In, Out)>,
184}
185
186impl<F, Sf, St, Req, In, Out, Err> ApplyFactory<F, Sf, St, Req, In, Out, Err>
187where
188 F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
189 Sf: ServiceFactory<St, Req>,
190{
191 pub(crate) fn new(sf: Sf, f: F) -> Self
193 where
194 Sf: ServiceFactory<St, Req>,
195 Err: From<Sf::Error>,
196 {
197 Self {
198 f,
199 sf,
200 r: marker::PhantomData,
201 }
202 }
203}
204
205impl<F, Sf, St, Req, In, Out, Err> Clone for ApplyFactory<F, Sf, St, Req, In, Out, Err>
206where
207 F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
208 Sf: ServiceFactory<St, Req> + Clone,
209{
210 fn clone(&self) -> Self {
211 Self {
212 f: self.f.clone(),
213 sf: self.sf.clone(),
214 r: marker::PhantomData,
215 }
216 }
217}
218
219impl<F, Sf, St, Req, In, Out, Err> fmt::Debug for ApplyFactory<F, Sf, St, Req, In, Out, Err>
220where
221 F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
222 Sf: ServiceFactory<St, Req> + fmt::Debug,
223{
224 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
225 f.debug_struct("ApplyFactory")
226 .field("factory", &self.sf)
227 .field("map", &std::any::type_name::<F>())
228 .finish()
229 }
230}
231
232impl<F, Sf, St, Req, In, Out, Err> ServiceFactory<St, In>
233 for ApplyFactory<F, Sf, St, Req, In, Out, Err>
234where
235 F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
236 Sf: ServiceFactory<St, Req>,
237 Err: From<Sf::Error>,
238{
239 type Res = Out;
240 type Error = Err;
241
242 type Service = Apply<Sf::Service, St, Req, F, In, Out, Err>;
243 type InitError = Sf::InitError;
244
245 #[inline]
246 async fn create(&self, st: &St) -> Result<Self::Service, Self::InitError> {
247 self.sf.create(st).await.map(|svc| Apply {
248 svc,
249 f: self.f.clone(),
250 r: marker::PhantomData,
251 ready: Cell::new(false),
252 entered: Cell::new(0),
253 })
254 }
255}
256
257#[cfg(test)]
258mod tests {
259 use std::{cell::Cell, rc::Rc};
260
261 use super::*;
262 use crate::{factory, fn_factory, service};
263
264 #[derive(Debug, Default, Clone)]
265 struct Srv(Rc<Cell<usize>>);
266
267 impl Service<(), ()> for Srv {
268 type Res = ();
269 type Error = ();
270
271 async fn call(&self, _r: (), _: Ctx<'_, Self>) -> Result<(), ()> {
272 Ok(())
273 }
274
275 async fn ready(&self, _: Ctx<'_, Self>) -> Result<(), ()> {
276 self.0.set(self.0.get() + 1);
277 Ok(())
278 }
279
280 async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
281 self.0.set(self.0.get() + 1);
282 }
283 }
284
285 #[derive(Debug, PartialEq, Eq)]
286 struct Err;
287
288 impl From<()> for Err {
289 fn from(_e: ()) -> Self {
290 Err
291 }
292 }
293
294 #[ntex::test]
295 async fn test_call() {
296 let cnt_sht = Rc::new(Cell::new(0));
297 let srv = service(
298 apply_fn(Srv(cnt_sht.clone()), async move |req: &'static str, svc| {
299 svc.call(()).await.unwrap();
300 Ok((req, ()))
301 })
302 .clone(),
303 )
304 .pipeline(());
305
306 assert_eq!(srv.ready().await, Ok::<_, Err>(()));
307
308 srv.shutdown().await;
309 assert_eq!(cnt_sht.get(), 2);
310
311 let res = srv.call("srv").await;
312 assert!(res.is_ok());
313 assert_eq!(res.unwrap(), ("srv", ()));
314 }
315
316 #[ntex::test]
317 async fn test_call_svc() {
318 let cnt_sht = Rc::new(Cell::new(0));
319 let srv = service(Srv(cnt_sht.clone()))
320 .apply_fn(async move |req: &'static str, svc| {
321 svc.st();
322 svc.call(()).await.unwrap();
323 Ok((req, ()))
324 })
325 .clone();
326 let s = format!("{srv:?}");
327 assert!(s.contains("Apply"), "{}", s);
328
329 let srv = srv.pipeline(());
330 assert_eq!(srv.ready().await, Ok::<_, Err>(()));
331
332 srv.shutdown().await;
333 assert_eq!(cnt_sht.get(), 2);
334
335 let res = srv.call("srv").await;
336 assert!(res.is_ok());
337 assert_eq!(res.unwrap(), ("srv", ()));
338 let _ = format!("{srv:?}");
339 }
340
341 #[ntex::test]
342 async fn test_create() {
343 let new_srv = factory(apply_fn_factory(
344 fn_factory(|(): &()| async { Ok::<_, ()>(Srv::default()) }),
345 async move |req: &'static str, srv| {
346 srv.call(()).await.unwrap();
347 Ok((req, ()))
348 },
349 ));
350
351 let srv = new_srv.pipeline(()).await.unwrap();
352
353 assert_eq!(srv.ready().await, Ok::<_, Err>(()));
354
355 let res = srv.call("srv").await;
356 assert!(res.is_ok());
357 assert_eq!(res.unwrap(), ("srv", ()));
358 assert_eq!(Err, Err::from(()));
359 }
360
361 #[ntex::test]
362 async fn test_create_chain() {
363 let new_srv = factory(fn_factory(|(): &()| async { Ok::<_, ()>(Srv::default()) }))
364 .apply_fn(async move |req: &'static str, srv| {
365 srv.call(()).await.unwrap();
366 Ok((req, ()))
367 })
368 .clone();
369
370 let srv = new_srv.pipeline(()).await.unwrap();
371
372 assert_eq!(srv.ready().await, Ok::<_, Err>(()));
373
374 let res = srv.call("srv").await;
375 assert!(res.is_ok());
376 assert_eq!(res.unwrap(), ("srv", ()));
377 let _ = format!("{new_srv:?}");
378 }
379
380 mod ready_flag {
381 use std::rc::Rc;
382
383 use crate::util::tests::{Req, Single, State, concurrent_entry};
384 use crate::{Pipeline, apply::Apply};
385
386 #[ntex::test]
387 async fn concurrent_entry_after_await() {
388 let st = Rc::new(State::default());
389 let pl = Pipeline::new(
390 (),
391 Apply::new(Single(st.clone()), async |(rx1, rx2): Req, svc| {
392 let _ = rx1.await;
393 svc.call(rx2).await
394 }),
395 );
396 concurrent_entry(&pl, &st).await;
397 }
398 }
399}