Skip to main content

snafu/futures/
try_stream.rs

1//! Additions to the [`TryStream`] trait.
2//!
3//! [`TryStream`]: futures_core::TryStream
4
5use crate::IntoError;
6use core::{
7    marker::PhantomData,
8    pin::Pin,
9    task::{Context as TaskContext, Poll},
10};
11use futures_core::stream::{Stream, TryStream};
12use pin_project::pin_project;
13
14#[cfg(any(feature = "alloc", test))]
15use alloc::string::String;
16
17#[cfg(any(feature = "alloc", test))]
18use crate::FromString;
19
20/// Additions to [`TryStream`].
21pub trait TryStreamExt: TryStream + Sized {
22    /// Extend a [`TryStream`]'s error with additional context-sensitive
23    /// information.
24    ///
25    /// ```rust
26    /// # #[cfg(feature = "internal-dev-dependencies")] {
27    /// use futures::TryStream;
28    /// # use futures::stream;
29    /// use snafu::prelude::*;
30    ///
31    /// #[derive(Debug, Snafu)]
32    /// enum Error {
33    ///     Authenticating {
34    ///         user_name: String,
35    ///         user_id: i32,
36    ///         source: ApiError,
37    ///     },
38    /// }
39    ///
40    /// fn example() -> impl TryStream<Ok = i32, Error = Error> {
41    ///     stock_prices().context(AuthenticatingSnafu {
42    ///         user_name: "admin",
43    ///         user_id: 42,
44    ///     })
45    /// }
46    ///
47    /// # type ApiError = Box<dyn std::error::Error>;
48    /// fn stock_prices() -> impl TryStream<Ok = i32, Error = ApiError> {
49    ///     /* ... */
50    /// # stream::empty()
51    /// }
52    /// }
53    /// ```
54    ///
55    /// Note that the context selector will call [`Into::into`] on
56    /// each field, so the types are not required to exactly match.
57    fn context<C, E>(self, context: C) -> Context<Self, C, E>
58    where
59        C: IntoError<E, Source = Self::Error> + Clone;
60
61    /// Extend a [`TryStream`]'s error with lazily-generated
62    /// context-sensitive information.
63    ///
64    /// ```rust
65    /// # #[cfg(feature = "internal-dev-dependencies")] {
66    /// use futures::TryStream;
67    /// # use futures::stream;
68    /// use snafu::prelude::*;
69    ///
70    /// #[derive(Debug, Snafu)]
71    /// enum Error {
72    ///     Authenticating {
73    ///         user_name: String,
74    ///         user_id: i32,
75    ///         source: ApiError,
76    ///     },
77    /// }
78    ///
79    /// fn example() -> impl TryStream<Ok = i32, Error = Error> {
80    ///     stock_prices().with_context(|_| AuthenticatingSnafu {
81    ///         user_name: "admin",
82    ///         user_id: 42,
83    ///     })
84    /// }
85    ///
86    /// # type ApiError = Box<dyn std::error::Error>;
87    /// fn stock_prices() -> impl TryStream<Ok = i32, Error = ApiError> {
88    ///     /* ... */
89    /// # stream::empty()
90    /// }
91    /// # }
92    /// ```
93    ///
94    /// Note that this *may not* be needed in many cases because the
95    /// context selector will call [`Into::into`] on each field.
96    fn with_context<F, C, E>(self, context: F) -> WithContext<Self, F, E>
97    where
98        F: FnMut(&mut Self::Error) -> C,
99        C: IntoError<E, Source = Self::Error>;
100
101    /// Extend a [`TryStream`]'s error with information from a string.
102    ///
103    /// The target error type must implement [`FromString`] by using
104    /// the
105    /// [`#[snafu(whatever)]`][crate::Snafu#controlling-stringly-typed-errors]
106    /// attribute. The premade [`Whatever`](crate::Whatever) type is also available.
107    ///
108    /// In many cases, you will want to use
109    /// [`with_whatever_context`][Self::with_whatever_context] instead
110    /// as it is only called in case of error. This method is best
111    /// suited for when you have a string literal.
112    ///
113    /// ```rust
114    /// # #[cfg(feature = "internal-dev-dependencies")] {
115    /// use futures::TryStream;
116    /// # use futures::stream;
117    /// use snafu::{prelude::*, Whatever};
118    ///
119    /// fn example() -> impl TryStream<Ok = i32, Error = Whatever> {
120    ///     stock_prices().whatever_context("Couldn't get stock prices")
121    /// }
122    ///
123    /// # type ApiError = Box<dyn std::error::Error + Send + Sync>;
124    /// fn stock_prices() -> impl TryStream<Ok = i32, Error = ApiError> {
125    ///     /* ... */
126    /// # stream::empty()
127    /// }
128    /// # }
129    /// ```
130    #[cfg(any(feature = "alloc", test))]
131    fn whatever_context<S, E>(self, context: S) -> WhateverContext<Self, S, E>
132    where
133        S: Into<String>,
134        E: FromString;
135
136    /// Extend a [`TryStream`]'s error with information from a
137    /// lazily-generated string.
138    ///
139    /// The target error type must implement [`FromString`] by using
140    /// the
141    /// [`#[snafu(whatever)]`][crate::Snafu#controlling-stringly-typed-errors]
142    /// attribute. The premade [`Whatever`](crate::Whatever) type is also available.
143    ///
144    /// ```rust
145    /// # #[cfg(feature = "internal-dev-dependencies")] {
146    /// use futures::TryStream;
147    /// # use futures::stream;
148    /// use snafu::{prelude::*, Whatever};
149    ///
150    /// fn example(symbol: &'static str) -> impl TryStream<Ok = i32, Error = Whatever> {
151    ///     stock_prices(symbol)
152    ///         .with_whatever_context(move |_| format!("Couldn't get stock prices for {symbol}"))
153    /// }
154    ///
155    /// # type ApiError = Box<dyn std::error::Error + Send + Sync>;
156    /// fn stock_prices(symbol: &'static str) -> impl TryStream<Ok = i32, Error = ApiError> {
157    ///     /* ... */
158    /// # stream::empty()
159    /// }
160    /// # }
161    /// ```
162    #[cfg(any(feature = "alloc", test))]
163    fn with_whatever_context<F, S, E>(self, context: F) -> WithWhateverContext<Self, F, E>
164    where
165        F: FnMut(&mut Self::Error) -> S,
166        S: Into<String>,
167        E: FromString;
168}
169
170impl<St> TryStreamExt for St
171where
172    St: TryStream,
173{
174    fn context<C, E>(self, context: C) -> Context<Self, C, E>
175    where
176        C: IntoError<E, Source = Self::Error> + Clone,
177    {
178        Context {
179            inner: self,
180            context,
181            _e: PhantomData,
182        }
183    }
184
185    fn with_context<F, C, E>(self, context: F) -> WithContext<Self, F, E>
186    where
187        F: FnMut(&mut Self::Error) -> C,
188        C: IntoError<E, Source = Self::Error>,
189    {
190        WithContext {
191            inner: self,
192            context,
193            _e: PhantomData,
194        }
195    }
196
197    #[cfg(any(feature = "alloc", test))]
198    fn whatever_context<S, E>(self, context: S) -> WhateverContext<Self, S, E>
199    where
200        S: Into<String>,
201        E: FromString,
202    {
203        WhateverContext {
204            inner: self,
205            context,
206            _e: PhantomData,
207        }
208    }
209
210    #[cfg(any(feature = "alloc", test))]
211    fn with_whatever_context<F, S, E>(self, context: F) -> WithWhateverContext<Self, F, E>
212    where
213        F: FnMut(&mut Self::Error) -> S,
214        S: Into<String>,
215        E: FromString,
216    {
217        WithWhateverContext {
218            inner: self,
219            context,
220            _e: PhantomData,
221        }
222    }
223}
224
225/// Stream for the [`context`](TryStreamExt::context) combinator.
226///
227/// See the [`TryStreamExt::context`] method for more details.
228#[pin_project]
229#[derive(Debug)]
230#[must_use = "streams do nothing unless polled"]
231pub struct Context<St, C, E> {
232    #[pin]
233    inner: St,
234    context: C,
235    _e: PhantomData<E>,
236}
237
238impl<St, C, E> Stream for Context<St, C, E>
239where
240    St: TryStream,
241    C: IntoError<E, Source = St::Error> + Clone,
242{
243    type Item = Result<St::Ok, E>;
244
245    #[track_caller]
246    fn poll_next(self: Pin<&mut Self>, ctx: &mut TaskContext) -> Poll<Option<Self::Item>> {
247        let this = self.project();
248        let inner = this.inner;
249        let context = this.context;
250
251        match inner.try_poll_next(ctx) {
252            Poll::Pending => Poll::Pending,
253            Poll::Ready(None) => Poll::Ready(None),
254            Poll::Ready(Some(Ok(v))) => Poll::Ready(Some(Ok(v))),
255            Poll::Ready(Some(Err(error))) => {
256                let error = context.clone().into_error(error);
257                Poll::Ready(Some(Err(error)))
258            }
259        }
260    }
261}
262
263/// Stream for the [`with_context`](TryStreamExt::with_context) combinator.
264///
265/// See the [`TryStreamExt::with_context`] method for more details.
266#[pin_project]
267#[derive(Debug)]
268#[must_use = "streams do nothing unless polled"]
269pub struct WithContext<St, F, E> {
270    #[pin]
271    inner: St,
272    context: F,
273    _e: PhantomData<E>,
274}
275
276impl<St, F, C, E> Stream for WithContext<St, F, E>
277where
278    St: TryStream,
279    F: FnMut(&mut St::Error) -> C,
280    C: IntoError<E, Source = St::Error>,
281{
282    type Item = Result<St::Ok, E>;
283
284    #[track_caller]
285    fn poll_next(self: Pin<&mut Self>, ctx: &mut TaskContext) -> Poll<Option<Self::Item>> {
286        let this = self.project();
287        let inner = this.inner;
288        let context = this.context;
289
290        match inner.try_poll_next(ctx) {
291            Poll::Pending => Poll::Pending,
292            Poll::Ready(None) => Poll::Ready(None),
293            Poll::Ready(Some(Ok(v))) => Poll::Ready(Some(Ok(v))),
294            Poll::Ready(Some(Err(mut error))) => {
295                let error = context(&mut error).into_error(error);
296                Poll::Ready(Some(Err(error)))
297            }
298        }
299    }
300}
301
302/// Stream for the
303/// [`whatever_context`](TryStreamExt::whatever_context) combinator.
304///
305/// See the [`TryStreamExt::whatever_context`] method for more
306/// details.
307#[pin_project]
308#[derive(Debug)]
309#[must_use = "streams do nothing unless polled"]
310#[cfg(any(feature = "alloc", test))]
311pub struct WhateverContext<St, S, E> {
312    #[pin]
313    inner: St,
314    context: S,
315    _e: PhantomData<E>,
316}
317
318#[cfg(any(feature = "alloc", test))]
319impl<St, S, E> Stream for WhateverContext<St, S, E>
320where
321    St: TryStream,
322    S: Into<String> + Clone,
323    E: FromString,
324    St::Error: Into<E::Source>,
325{
326    type Item = Result<St::Ok, E>;
327
328    #[track_caller]
329    fn poll_next(self: Pin<&mut Self>, ctx: &mut TaskContext) -> Poll<Option<Self::Item>> {
330        let this = self.project();
331        let inner = this.inner;
332        let context = this.context;
333
334        match inner.try_poll_next(ctx) {
335            Poll::Pending => Poll::Pending,
336            Poll::Ready(None) => Poll::Ready(None),
337            Poll::Ready(Some(Ok(v))) => Poll::Ready(Some(Ok(v))),
338            Poll::Ready(Some(Err(error))) => {
339                let error = E::with_source(error.into(), context.clone().into());
340                Poll::Ready(Some(Err(error)))
341            }
342        }
343    }
344}
345
346/// Stream for the
347/// [`with_whatever_context`](TryStreamExt::with_whatever_context)
348/// combinator.
349///
350/// See the [`TryStreamExt::with_whatever_context`] method for more
351/// details.
352#[pin_project]
353#[derive(Debug)]
354#[must_use = "streams do nothing unless polled"]
355#[cfg(any(feature = "alloc", test))]
356pub struct WithWhateverContext<St, F, E> {
357    #[pin]
358    inner: St,
359    context: F,
360    _e: PhantomData<E>,
361}
362
363#[cfg(any(feature = "alloc", test))]
364impl<St, F, S, E> Stream for WithWhateverContext<St, F, E>
365where
366    St: TryStream,
367    F: FnMut(&mut St::Error) -> S,
368    S: Into<String>,
369    E: FromString,
370    St::Error: Into<E::Source>,
371{
372    type Item = Result<St::Ok, E>;
373
374    #[track_caller]
375    fn poll_next(self: Pin<&mut Self>, ctx: &mut TaskContext) -> Poll<Option<Self::Item>> {
376        let this = self.project();
377        let inner = this.inner;
378        let context = this.context;
379
380        match inner.try_poll_next(ctx) {
381            Poll::Pending => Poll::Pending,
382            Poll::Ready(None) => Poll::Ready(None),
383            Poll::Ready(Some(Ok(v))) => Poll::Ready(Some(Ok(v))),
384            Poll::Ready(Some(Err(mut error))) => {
385                let context = context(&mut error);
386                let error = E::with_source(error.into(), context.into());
387                Poll::Ready(Some(Err(error)))
388            }
389        }
390    }
391}