Skip to main content

oas3_gen_support/
event_stream.rs

1use 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
20/// A stream of Server-Sent Events (SSE) that deserializes each event's data as JSON.
21///
22/// This wraps a `reqwest::Response` and parses the SSE event stream, deserializing
23/// each event's `data` field as the type parameter `T`.
24///
25/// # Example
26///
27/// ```ignore
28/// use futures::StreamExt;
29///
30/// let response = client.get("/events").send().await?;
31/// let mut stream = EventStream::<MyEvent>::from_response(response);
32///
33/// while let Some(result) = stream.next().await {
34///     match result {
35///         Ok(event) => println!("Received: {:?}", event),
36///         Err(e) => eprintln!("Error: {}", e),
37///     }
38/// }
39/// ```
40pub 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  /// Create an `EventStream` from an HTTP response.
60  ///
61  /// The response should have content type `text/event-stream`.
62  #[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}