summaryrefslogtreecommitdiffstats
path: root/src/client/ernie.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-29 08:43:23 +0800
committerGitHub <noreply@github.com>2024-04-29 08:43:23 +0800
commit68882ecd4dced38e92b6ef581a30e45358fc61e0 (patch)
tree356a3de5ff8d3a0cb7defea1ad91924894f2f9df /src/client/ernie.rs
parent37a0cd08a92f07ef24e39bdab7c8bced5b59c146 (diff)
downloadaichat-68882ecd4dced38e92b6ef581a30e45358fc61e0.tar.gz
refactor: abstract event stream handling (#458)
Diffstat (limited to 'src/client/ernie.rs')
-rw-r--r--src/client/ernie.rs56
1 files changed, 10 insertions, 46 deletions
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index cbfdf22..4038e0d 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -1,16 +1,14 @@
use super::{
- maybe_catch_error, patch_system_message, Client, CompletionDetails, ErnieClient, ExtraConfig,
- Model, ModelConfig, PromptType, SendData, SseHandler,
+ maybe_catch_error, patch_system_message, sse_stream, Client, CompletionDetails, ErnieClient,
+ ExtraConfig, Model, ModelConfig, PromptType, SendData, SseHandler,
};
use crate::utils::PromptKind;
-use anyhow::{anyhow, bail, Context, Result};
+use anyhow::{anyhow, Context, Result};
use async_trait::async_trait;
use chrono::Utc;
-use futures_util::StreamExt;
use reqwest::{Client as ReqwestClient, RequestBuilder};
-use reqwest_eventsource::{Error as EventSourceError, Event, RequestBuilderExt};
use serde::Deserialize;
use serde_json::{json, Value};
use std::env;
@@ -108,49 +106,15 @@ async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionDeta
}
async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandler) -> Result<()> {
- let mut es = builder.eventsource()?;
- while let Some(event) = es.next().await {
- match event {
- Ok(Event::Open) => {}
- Ok(Event::Message(message)) => {
- let data: Value = serde_json::from_str(&message.data)?;
- if let Some(text) = data["result"].as_str() {
- handler.text(text)?;
- }
- }
- Err(err) => {
- match err {
- EventSourceError::InvalidContentType(header_value, res) => {
- let content_type = header_value
- .to_str()
- .map_err(|_| anyhow!("Invalid response header"))?;
- if content_type.contains("application/json") {
- let data: Value = res.json().await?;
- maybe_catch_error(&data)?;
- bail!("Invalid response data: {data}");
- } else {
- let text = res.text().await?;
- if let Some(text) = text.strip_prefix("data: ") {
- let data: Value = serde_json::from_str(text)?;
- if let Some(text) = data["result"].as_str() {
- handler.text(text)?;
- }
- } else {
- bail!("Invalid response data: {text}")
- }
- }
- }
- EventSourceError::StreamEnded => {}
- _ => {
- bail!("{}", err);
- }
- }
- es.close();
- }
+ let handle = |data: &str| -> Result<bool> {
+ let data: Value = serde_json::from_str(data)?;
+ if let Some(text) = data["result"].as_str() {
+ handler.text(text)?;
}
- }
+ Ok(false)
+ };
- Ok(())
+ sse_stream(builder, handle).await
}
fn build_body(data: SendData, model: &Model) -> Value {