Skip to main content

xmtp_api_grpc/streams/
default.rs

1//! Default XMTP Streams
2
3use prost::bytes::Bytes;
4use std::{
5    marker::PhantomData,
6    pin::Pin,
7    task::{Context, Poll, ready},
8};
9
10use crate::error::GrpcError;
11use futures::{Stream, TryStream};
12use pin_project::pin_project;
13use xmtp_proto::{ApiEndpoint, api::ApiClientError};
14
15#[pin_project]
16/// A stream which maps the tonic error to ApiClientError, and attaches endpoint metadata
17pub struct XmtpTonicStream<S, T> {
18    #[pin]
19    inner: S,
20    endpoint: ApiEndpoint,
21    _marker: PhantomData<T>,
22}
23
24impl<S, T> XmtpTonicStream<S, T> {
25    pub fn new(inner: S, endpoint: ApiEndpoint) -> Self {
26        Self {
27            inner,
28            endpoint,
29            _marker: PhantomData,
30        }
31    }
32}
33
34impl<S, T> Stream for XmtpTonicStream<S, T>
35where
36    S: TryStream<Ok = Bytes, Error = GrpcError>,
37    GrpcError: From<<S as TryStream>::Error>,
38    T: prost::Message + Default,
39{
40    type Item = Result<T, ApiClientError>;
41
42    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
43        let this = self.as_mut().project();
44        if let Some(item) = ready!(this.inner.try_poll_next(cx)) {
45            let res = item
46                .map_err(|e| ApiClientError::new(self.endpoint.clone(), e))
47                .and_then(|i| T::decode(i).map_err(GrpcError::from).map_err(Into::into));
48            Poll::Ready(Some(res))
49        } else {
50            Poll::Ready(None)
51        }
52    }
53}
54
55#[cfg(test)]
56mod tests {
57    use super::*;
58    use futures::{StreamExt, stream};
59    use prost::Message;
60    use rstest::rstest;
61
62    #[derive(Clone, PartialEq, Message)]
63    struct TestMessage {
64        #[prost(string, tag = "1")]
65        pub content: String,
66    }
67
68    impl prost::Name for TestMessage {
69        const NAME: &'static str = "TestMessage";
70        const PACKAGE: &'static str = "test";
71        fn full_name() -> String {
72            format!("{}.{}", Self::PACKAGE, Self::NAME)
73        }
74    }
75
76    fn create_test_message_bytes(content: &str) -> Bytes {
77        let msg = TestMessage {
78            content: content.to_string(),
79        };
80        Bytes::from(msg.encode_to_vec())
81    }
82
83    #[rstest]
84    #[case::empty_stream(vec![], vec![])]
85    #[case::single_message(
86        vec![Ok(create_test_message_bytes("test1"))],
87        vec![TestMessage { content: "test1".to_string() }],
88    )]
89    #[case::multiple_messages(
90        vec![
91            Ok(create_test_message_bytes("msg1")),
92            Ok(create_test_message_bytes("msg2")),
93            Ok(create_test_message_bytes("msg3"))
94        ],
95        vec![
96            TestMessage { content: "msg1".to_string() },
97            TestMessage { content: "msg2".to_string() },
98            TestMessage { content: "msg3".to_string() }
99        ],
100    )]
101    #[xmtp_common::test]
102    async fn test_successful_message_decoding(
103        #[case] input: Vec<Result<Bytes, GrpcError>>,
104        #[case] expected: Vec<TestMessage>,
105    ) {
106        let stream = stream::iter(input);
107        let endpoint = ApiEndpoint::SubscribeGroupMessages;
108        let stream = XmtpTonicStream::<_, TestMessage>::new(stream, endpoint);
109
110        let results: Vec<_> = stream.map(Result::unwrap).collect().await;
111        assert_eq!(results, expected);
112    }
113
114    #[xmtp_common::test]
115    async fn test_error_propagation() {
116        let grpc_error = GrpcError::Status(tonic::Status::unavailable("Connection failed"));
117        let input = vec![
118            Ok(create_test_message_bytes("msg1")),
119            Err(grpc_error),
120            Ok(create_test_message_bytes("msg3")),
121        ];
122
123        let stream = stream::iter(input);
124        let endpoint = ApiEndpoint::QueryGroupMessages;
125        let stream = XmtpTonicStream::<_, TestMessage>::new(stream, endpoint.clone());
126
127        let results: Vec<_> = stream.collect().await;
128        assert_eq!(results.len(), 3);
129
130        assert_eq!(
131            results[0].as_ref().unwrap(),
132            &TestMessage {
133                content: "msg1".to_string()
134            }
135        );
136
137        let api_error = results[1].as_ref().unwrap_err();
138        if let xmtp_proto::api::ApiClientError::ClientWithEndpoint {
139            endpoint: err_endpoint,
140            ..
141        } = api_error
142        {
143            assert_eq!(*err_endpoint, endpoint.to_string());
144        } else {
145            panic!("Expected ClientWithEndpoint error variant");
146        }
147
148        assert_eq!(
149            results[2].as_ref().unwrap(),
150            &TestMessage {
151                content: "msg3".to_string()
152            }
153        );
154    }
155
156    #[xmtp_common::test]
157    fn stream_ends() {
158        let input = vec![Ok(create_test_message_bytes("test"))];
159        let stream = stream::iter(input);
160        let endpoint = ApiEndpoint::SendGroupMessages;
161        let stream = XmtpTonicStream::<_, TestMessage>::new(stream, endpoint);
162
163        futures::pin_mut!(stream);
164        let mut cx = futures_test::task::noop_context();
165
166        let first_poll = stream.as_mut().poll_next(&mut cx);
167        assert!(matches!(first_poll, Poll::Ready(Some(Ok(_)))));
168        if let Poll::Ready(Some(Ok(msg))) = first_poll {
169            assert_eq!(msg.content, "test");
170        }
171
172        let end_poll = stream.poll_next(&mut cx);
173        assert!(matches!(end_poll, Poll::Ready(None)));
174    }
175}