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/ernie.rs | |
| parent | 37a0cd08a92f07ef24e39bdab7c8bced5b59c146 (diff) | |
| download | aichat-68882ecd4dced38e92b6ef581a30e45358fc61e0.tar.gz | |
refactor: abstract event stream handling (#458)
Diffstat (limited to 'src/client/ernie.rs')
| -rw-r--r-- | src/client/ernie.rs | 56 |
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 { |
