use crate::error::{Error, Result};
use crate::protocol_generated::types::Event;
use futures_core::Stream;
use reqwest::{Client, RequestBuilder};
use reqwest_eventsource::retry::ExponentialBackoff;
use reqwest_eventsource::{Event as EsEvent, EventSource, ReadyState};
use std::pin::Pin;
use std::task::{Context, Poll};
use std::time::Duration;
pub const EVENT_PATH: &str = "/event";
pub const GLOBAL_EVENT_PATH: &str = "/global/event";
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct RetryConfig {
pub initial_interval: Duration,
pub max_interval: Duration,
pub factor: f64,
pub max_retries: Option<usize>,
}
impl Default for RetryConfig {
fn default() -> Self {
Self {
initial_interval: Duration::from_millis(500),
max_interval: Duration::from_secs(30),
factor: 2.0,
max_retries: None,
}
}
}
impl RetryConfig {
fn into_policy(self) -> ExponentialBackoff {
ExponentialBackoff::new(
self.initial_interval,
self.factor,
Some(self.max_interval),
self.max_retries,
)
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum StreamEvent {
Connected,
Event(Box<Event>),
Unknown(serde_json::Value),
}
impl StreamEvent {
#[must_use]
pub fn as_event(&self) -> Option<&Event> {
match self {
Self::Event(ev) => Some(ev.as_ref()),
_ => None,
}
}
#[must_use]
pub fn is_connected(&self) -> bool {
matches!(self, Self::Connected)
}
}
pub struct EventStream {
inner: EventSource,
}
impl EventStream {
pub fn connect(base_url: &str) -> Result<Self> {
Self::connect_with(
&Client::builder().build()?,
base_url,
RetryConfig::default(),
)
}
pub fn connect_with(client: &Client, base_url: &str, retry: RetryConfig) -> Result<Self> {
let url = format!("{}{EVENT_PATH}", base_url.trim_end_matches('/'));
Self::from_request(client.get(url), retry)
}
pub fn from_request(builder: RequestBuilder, retry: RetryConfig) -> Result<Self> {
let mut source = EventSource::new(builder).map_err(|e| Error::Http {
status: 0,
body: format!("cannot prepare SSE request: {e}"),
})?;
source.set_retry_policy(Box::new(retry.into_policy()));
Ok(Self { inner: source })
}
pub async fn next(&mut self) -> Option<Result<StreamEvent>> {
std::future::poll_fn(|cx| Pin::new(&mut *self).poll_next(cx)).await
}
#[must_use]
pub fn ready_state(&self) -> ReadyState {
self.inner.ready_state()
}
pub fn close(&mut self) {
self.inner.close();
}
}
fn decode_frame(data: &str) -> Result<StreamEvent> {
let value: serde_json::Value = serde_json::from_str(data)?;
match serde_json::from_value::<Event>(value.clone()) {
Ok(event) => Ok(StreamEvent::Event(Box::new(event))),
Err(_) => Ok(StreamEvent::Unknown(value)),
}
}
impl Stream for EventStream {
type Item = Result<StreamEvent>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
loop {
match Pin::new(&mut this.inner).poll_next(cx) {
Poll::Ready(Some(Ok(EsEvent::Open))) => {
return Poll::Ready(Some(Ok(StreamEvent::Connected)));
}
Poll::Ready(Some(Ok(EsEvent::Message(message)))) => {
return Poll::Ready(Some(decode_frame(&message.data)));
}
Poll::Ready(Some(Err(e))) => {
if this.inner.ready_state() == ReadyState::Closed {
return Poll::Ready(Some(Err(e.into())));
}
continue;
}
Poll::Ready(None) => return Poll::Ready(None),
Poll::Pending => return Poll::Pending,
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn retry_config_default_is_capped_and_unbounded() {
let cfg = RetryConfig::default();
assert_eq!(cfg.initial_interval, Duration::from_millis(500));
assert_eq!(cfg.max_interval, Duration::from_secs(30));
assert_eq!(cfg.max_retries, None);
}
#[test]
fn decode_known_event_frame() {
let frame = r#"{"type":"session.idle","id":"evt_1","properties":{"sessionID":"ses_123"}}"#;
let decoded = decode_frame(frame).expect("valid json");
match decoded {
StreamEvent::Event(ev) if matches!(*ev, Event::SessionIdle(_)) => {}
other => panic!("expected SessionIdle event, got {other:?}"),
}
}
#[test]
fn decode_unknown_type_becomes_unknown_variant() {
let frame = r#"{"type":"totally.new.event.from.the.future","properties":{"x":1}}"#;
let decoded = decode_frame(frame).expect("valid json");
match decoded {
StreamEvent::Unknown(value) => {
assert_eq!(value["type"], "totally.new.event.from.the.future");
}
other => panic!("expected Unknown, got {other:?}"),
}
}
#[test]
fn decode_known_type_wrong_shape_becomes_unknown_not_error() {
let frame = r#"{"type":"session.idle","id":"evt_1","properties":{"sessionID":42}}"#;
let decoded = decode_frame(frame).expect("valid json");
assert!(matches!(decoded, StreamEvent::Unknown(_)));
}
#[test]
fn decode_non_json_frame_is_error() {
assert!(decode_frame("not json at all").is_err());
}
#[test]
fn stream_event_accessors() {
assert!(StreamEvent::Connected.is_connected());
assert!(StreamEvent::Connected.as_event().is_none());
let ev =
decode_frame(r#"{"type":"session.idle","id":"evt_1","properties":{"sessionID":"s"}}"#)
.unwrap();
assert!(ev.as_event().is_some());
assert!(!ev.is_connected());
}
#[cfg(feature = "integration-tests")]
mod live {
use super::*;
const BASE_URL: &str = "http://127.0.0.1:41999";
async fn next_item(stream: &mut EventStream) -> Option<Result<StreamEvent>> {
stream.next().await
}
#[tokio::test]
async fn connects_and_receives_real_frames() {
let client = Client::new();
let mut stream =
EventStream::connect_with(&client, BASE_URL, RetryConfig::default()).unwrap();
let first = tokio::time::timeout(Duration::from_secs(5), next_item(&mut stream))
.await
.expect("server did not open the SSE stream in time")
.expect("stream closed immediately")
.expect("connection error");
assert!(first.is_connected(), "expected Connected, got {first:?}");
let created = client
.post(format!("{BASE_URL}/session"))
.json(&serde_json::json!({}))
.send()
.await
.expect("create session request failed");
assert!(
created.status().is_success(),
"POST /session failed: {}",
created.status()
);
let mut saw_event = false;
for _ in 0..50 {
match tokio::time::timeout(Duration::from_secs(10), next_item(&mut stream)).await {
Ok(Some(Ok(StreamEvent::Event(_) | StreamEvent::Unknown(_)))) => {
saw_event = true;
break;
}
Ok(Some(Ok(StreamEvent::Connected))) => continue,
Ok(Some(Err(e))) => panic!("stream error: {e}"),
Ok(None) => panic!("stream closed unexpectedly"),
Err(_) => panic!("timed out waiting for an event frame"),
}
}
assert!(saw_event, "no event frames observed after POST /session");
}
#[tokio::test]
async fn event_stream_is_a_futures_stream() {
fn assert_stream<S: Stream>(_: &S) {}
let stream = EventStream::connect(BASE_URL).unwrap();
assert_stream(&stream);
}
}
}