summaryrefslogtreecommitdiffstats
path: root/src/client/common.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/common.rs')
-rw-r--r--src/client/common.rs48
1 files changed, 48 insertions, 0 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index 64e32ff..e35e956 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -11,6 +11,7 @@ use async_trait::async_trait;
use futures_util::{Stream, StreamExt};
use lazy_static::lazy_static;
use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy, RequestBuilder};
+use reqwest_eventsource::{Error as EventSourceError, Event, RequestBuilderExt};
use serde::Deserialize;
use serde_json::{json, Value};
use std::{env, future::Future, time::Duration};
@@ -531,6 +532,53 @@ pub fn maybe_catch_error(data: &Value) -> Result<()> {
Ok(())
}
+pub async fn sse_stream<F>(builder: RequestBuilder, mut handle: F) -> Result<()>
+where
+ F: FnMut(&str) -> Result<bool>,
+{
+ let mut es = builder.eventsource()?;
+ while let Some(event) = es.next().await {
+ match event {
+ Ok(Event::Open) => {}
+ Ok(Event::Message(message)) => {
+ if handle(&message.data)? {
+ break;
+ }
+ }
+ Err(err) => {
+ match err {
+ EventSourceError::StreamEnded => {}
+ EventSourceError::InvalidStatusCode(status, res) => {
+ let text = res.text().await?;
+ let data: Value = match text.parse() {
+ Ok(data) => data,
+ Err(_) => {
+ bail!(
+ "Invalid response data: {text} (status: {})",
+ status.as_u16()
+ );
+ }
+ };
+ catch_error(&data, status.as_u16())?;
+ }
+ EventSourceError::InvalidContentType(header_value, res) => {
+ let text = res.text().await?;
+ bail!(
+ "Invalid response event-stream. content-type: {}, data: {text}",
+ header_value.to_str().unwrap_or_default()
+ );
+ }
+ _ => {
+ bail!("{}", err);
+ }
+ }
+ es.close();
+ }
+ }
+ }
+ Ok(())
+}
+
pub async fn json_stream<S, F>(mut stream: S, mut handle: F) -> Result<()>
where
S: Stream<Item = Result<bytes::Bytes, reqwest::Error>> + Unpin,