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/openai.rs | |
| parent | 37a0cd08a92f07ef24e39bdab7c8bced5b59c146 (diff) | |
| download | aichat-68882ecd4dced38e92b6ef581a30e45358fc61e0.tar.gz | |
refactor: abstract event stream handling (#458)
Diffstat (limited to 'src/client/openai.rs')
| -rw-r--r-- | src/client/openai.rs | 59 |
1 files changed, 13 insertions, 46 deletions
diff --git a/src/client/openai.rs b/src/client/openai.rs index 4c4eae2..ad833c7 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,14 +1,12 @@ use super::{ - catch_error, CompletionDetails, ExtraConfig, Model, ModelConfig, OpenAIClient, PromptType, - SendData, SseHandler, + catch_error, sse_stream, CompletionDetails, ExtraConfig, Model, ModelConfig, OpenAIClient, + PromptType, SendData, SseHandler, }; use crate::utils::PromptKind; -use anyhow::{anyhow, bail, Result}; -use futures_util::StreamExt; +use anyhow::{anyhow, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; -use reqwest_eventsource::{Error as EventSourceError, Event, RequestBuilderExt}; use serde::Deserialize; use serde_json::{json, Value}; @@ -67,49 +65,18 @@ pub async fn openai_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)) => { - if message.data == "[DONE]" { - break; - } - let data: Value = serde_json::from_str(&message.data)?; - if let Some(text) = data["choices"][0]["delta"]["content"].as_str() { - handler.text(text)?; - } - } - Err(err) => { - match err { - 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::StreamEnded => {} - EventSourceError::InvalidContentType(_, res) => { - let text = res.text().await?; - bail!("The API server should return data as 'text/event-stream', but it isn't. Check the client config. {text}"); - } - _ => { - bail!("{}", err); - } - } - es.close(); - } + let handle = |data: &str| -> Result<bool> { + if data == "[DONE]" { + return Ok(true); } - } + let data: Value = serde_json::from_str(data)?; + if let Some(text) = data["choices"][0]["delta"]["content"].as_str() { + handler.text(text)?; + } + Ok(false) + }; - Ok(()) + sse_stream(builder, handle).await } pub fn openai_build_body(data: SendData, model: &Model) -> Value { |
