xmtp_api_grpc/streams/
default.rs1use 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]
16pub 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}