oas3_gen_support/
event_stream.rs1use std::{
2 marker::PhantomData,
3 pin::Pin,
4 task::{Context, Poll},
5};
6
7use eventsource_stream::Eventsource;
8use futures_core::Stream;
9use serde::de::DeserializeOwned;
10
11#[derive(Debug, thiserror::Error)]
12pub enum EventStreamError {
13 #[error("SSE parse error: {0}")]
14 SseParse(#[from] eventsource_stream::EventStreamError<reqwest::Error>),
15
16 #[error("JSON deserialization error at path {path}: {inner}")]
17 JsonDeserialize { path: String, inner: serde_json::Error },
18}
19
20pub struct EventStream<T> {
41 inner: Pin<
42 Box<
43 dyn Stream<Item = Result<eventsource_stream::Event, eventsource_stream::EventStreamError<reqwest::Error>>> + Send,
44 >,
45 >,
46 _marker: PhantomData<T>,
47}
48
49impl<T> std::fmt::Debug for EventStream<T> {
50 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
51 f.debug_struct("EventStream").finish_non_exhaustive()
52 }
53}
54
55impl<T> EventStream<T>
56where
57 T: DeserializeOwned,
58{
59 #[must_use]
63 pub fn from_response(response: reqwest::Response) -> Self {
64 let stream = response.bytes_stream().eventsource();
65 Self {
66 inner: Box::pin(stream),
67 _marker: PhantomData,
68 }
69 }
70
71 fn parse_event(data: &str) -> Result<T, EventStreamError> {
72 let mut de = serde_json::Deserializer::from_str(data);
73 serde_path_to_error::deserialize(&mut de).map_err(|err| EventStreamError::JsonDeserialize {
74 path: err.path().to_string(),
75 inner: err.into_inner(),
76 })
77 }
78}
79
80impl<T> Stream for EventStream<T>
81where
82 T: DeserializeOwned + Unpin,
83{
84 type Item = Result<T, EventStreamError>;
85
86 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
87 loop {
88 match self.inner.as_mut().poll_next(cx) {
89 Poll::Ready(Some(event_result)) => match event_result {
90 Ok(event) => {
91 if event.data.is_empty() {
92 continue;
93 }
94 return Poll::Ready(Some(Self::parse_event(&event.data)));
95 }
96 Err(e) => return Poll::Ready(Some(Err(EventStreamError::SseParse(e)))),
97 },
98 Poll::Ready(None) => return Poll::Ready(None),
99 Poll::Pending => return Poll::Pending,
100 }
101 }
102 }
103}
104
105#[cfg(test)]
106mod tests {
107 use super::*;
108
109 #[derive(Debug, serde::Deserialize, PartialEq)]
110 struct TestEvent {
111 id: i32,
112 message: String,
113 }
114
115 #[test]
116 fn test_parse_event_success() {
117 let json = r#"{"id": 1, "message": "hello"}"#;
118 let result: Result<TestEvent, EventStreamError> = EventStream::<TestEvent>::parse_event(json);
119
120 assert!(result.is_ok());
121 let event = result.unwrap();
122 assert_eq!(event.id, 1);
123 assert_eq!(event.message, "hello");
124 }
125
126 #[test]
127 fn test_parse_event_invalid_json() {
128 let json = r#"{"id": "not_a_number", "message": "hello"}"#;
129 let result: Result<TestEvent, EventStreamError> = EventStream::<TestEvent>::parse_event(json);
130
131 assert!(result.is_err());
132 match result.unwrap_err() {
133 EventStreamError::JsonDeserialize { path, .. } => {
134 assert_eq!(path, "id");
135 }
136 EventStreamError::SseParse(err) => panic!("Expected JsonDeserialize error, got SseParse: {err}"),
137 }
138 }
139
140 #[test]
141 fn test_parse_event_missing_field() {
142 let json = r#"{"id": 1}"#;
143 let result: Result<TestEvent, EventStreamError> = EventStream::<TestEvent>::parse_event(json);
144
145 assert!(result.is_err());
146 match result.unwrap_err() {
147 EventStreamError::JsonDeserialize { path, .. } => {
148 assert!(path.contains("message") || path == ".");
149 }
150 EventStreamError::SseParse(err) => panic!("Expected JsonDeserialize error, got SseParse: {err}"),
151 }
152 }
153
154 #[test]
155 fn test_parse_event_empty_json() {
156 let json = r"{}";
157 let result: Result<TestEvent, EventStreamError> = EventStream::<TestEvent>::parse_event(json);
158
159 assert!(result.is_err());
160 }
161
162 #[derive(Debug, serde::Deserialize, PartialEq)]
163 struct NestedEvent {
164 data: InnerData,
165 }
166
167 #[derive(Debug, serde::Deserialize, PartialEq)]
168 struct InnerData {
169 value: String,
170 }
171
172 #[test]
173 fn test_parse_nested_event() {
174 let json = r#"{"data": {"value": "nested"}}"#;
175 let result: Result<NestedEvent, EventStreamError> = EventStream::<NestedEvent>::parse_event(json);
176
177 assert!(result.is_ok());
178 let event = result.unwrap();
179 assert_eq!(event.data.value, "nested");
180 }
181
182 #[test]
183 fn test_parse_nested_event_error_path() {
184 let json = r#"{"data": {"value": 123}}"#;
185 let result: Result<NestedEvent, EventStreamError> = EventStream::<NestedEvent>::parse_event(json);
186
187 assert!(result.is_err());
188 match result.unwrap_err() {
189 EventStreamError::JsonDeserialize { path, .. } => {
190 assert_eq!(path, "data.value");
191 }
192 EventStreamError::SseParse(err) => panic!("Expected JsonDeserialize error, got SseParse: {err}"),
193 }
194 }
195}