diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-29 08:43:23 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-29 08:43:23 +0800 |
| commit | 68882ecd4dced38e92b6ef581a30e45358fc61e0 (patch) | |
| tree | 356a3de5ff8d3a0cb7defea1ad91924894f2f9df /src/client/common.rs | |
| parent | 37a0cd08a92f07ef24e39bdab7c8bced5b59c146 (diff) | |
| download | aichat-68882ecd4dced38e92b6ef581a30e45358fc61e0.tar.gz | |
refactor: abstract event stream handling (#458)
Diffstat (limited to 'src/client/common.rs')
| -rw-r--r-- | src/client/common.rs | 48 |
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, |
