From 5458150ed3203cf13b0371efa2c791ac696cee93 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 23 May 2024 19:28:56 +0800 Subject: fix: json stream parser and refine client modules (#538) --- src/client/stream.rs | 292 +++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 292 insertions(+) create mode 100644 src/client/stream.rs (limited to 'src/client/stream.rs') diff --git a/src/client/stream.rs b/src/client/stream.rs new file mode 100644 index 0000000..7e2b5fa --- /dev/null +++ b/src/client/stream.rs @@ -0,0 +1,292 @@ +use super::{catch_error, ToolCall}; +use crate::utils::AbortSignal; + +use anyhow::{anyhow, bail, Context, Result}; +use futures_util::{Stream, StreamExt}; +use reqwest::RequestBuilder; +use reqwest_eventsource::{Error as EventSourceError, Event, RequestBuilderExt}; +use serde_json::Value; +use tokio::sync::mpsc::UnboundedSender; + +pub struct SseHandler { + sender: UnboundedSender, + abort: AbortSignal, + buffer: String, + tool_calls: Vec, +} + +impl SseHandler { + pub fn new(sender: UnboundedSender, abort: AbortSignal) -> Self { + Self { + sender, + abort, + buffer: String::new(), + tool_calls: Vec::new(), + } + } + + pub fn text(&mut self, text: &str) -> Result<()> { + // debug!("HandleText: {}", text); + if text.is_empty() { + return Ok(()); + } + self.buffer.push_str(text); + let ret = self + .sender + .send(SseEvent::Text(text.to_string())) + .with_context(|| "Failed to send ReplyEvent:Text"); + self.safe_ret(ret)?; + Ok(()) + } + + pub fn done(&mut self) -> Result<()> { + // debug!("HandleDone"); + let ret = self + .sender + .send(SseEvent::Done) + .with_context(|| "Failed to send ReplyEvent::Done"); + self.safe_ret(ret)?; + Ok(()) + } + + pub fn tool_call(&mut self, call: ToolCall) -> Result<()> { + // debug!("HandleCall: {:?}", call); + self.tool_calls.push(call); + Ok(()) + } + + pub fn get_abort(&self) -> AbortSignal { + self.abort.clone() + } + + pub fn take(self) -> (String, Vec) { + let Self { + buffer, tool_calls, .. + } = self; + (buffer, tool_calls) + } + + fn safe_ret(&self, ret: Result<()>) -> Result<()> { + if ret.is_err() && self.abort.aborted() { + return Ok(()); + } + ret + } +} + +#[derive(Debug)] +pub enum SseEvent { + Text(String), + Done, +} + +#[derive(Debug)] +pub struct SseMmessage { + pub event: String, + pub data: String, +} + +pub async fn sse_stream(builder: RequestBuilder, mut handle: F) -> Result<()> +where + F: FnMut(SseMmessage) -> Result, +{ + let mut es = builder.eventsource()?; + while let Some(event) = es.next().await { + match event { + Ok(Event::Open) => {} + Ok(Event::Message(message)) => { + let message = SseMmessage { + event: message.event, + data: message.data, + }; + if handle(message)? { + 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(mut stream: S, mut handle: F) -> Result<()> +where + S: Stream> + Unpin, + F: FnMut(&str) -> Result<()>, + E: std::error::Error, +{ + let mut parser = JsonStreamParser::default(); + let mut unparsed_bytes = vec![]; + while let Some(chunk_bytes) = stream.next().await { + let chunk_bytes = + chunk_bytes.map_err(|err| anyhow!("Failed to read json stream, {err}"))?; + unparsed_bytes.extend(chunk_bytes); + match std::str::from_utf8(&unparsed_bytes) { + Ok(text) => { + parser.process(text, &mut handle)?; + unparsed_bytes.clear(); + } + Err(_) => { + continue; + } + } + } + if !unparsed_bytes.is_empty() { + let text = std::str::from_utf8(&unparsed_bytes)?; + parser.process(text, &mut handle)?; + } + + Ok(()) +} + +#[derive(Debug, Default)] +struct JsonStreamParser { + buffer: Vec, + cursor: usize, + start: Option, + balances: Vec, + quoting: bool, + escape: bool, +} + +impl JsonStreamParser { + fn process(&mut self, text: &str, handle: &mut F) -> Result<()> + where + F: FnMut(&str) -> Result<()>, + { + self.buffer.extend(text.chars()); + + for i in self.cursor..self.buffer.len() { + let ch = self.buffer[i]; + if self.quoting { + if ch == '\\' { + self.escape = !self.escape; + } else { + if !self.escape && ch == '"' { + self.quoting = false; + } + self.escape = false; + } + continue; + } + match ch { + '"' => { + self.quoting = true; + self.escape = false; + } + '{' => { + if self.balances.is_empty() { + self.start = Some(i); + } + self.balances.push(ch); + } + '[' => { + if self.start.is_some() { + self.balances.push(ch); + } + } + '}' => { + self.balances.pop(); + if self.balances.is_empty() { + if let Some(start) = self.start.take() { + let value: String = self.buffer[start..=i].iter().collect(); + handle(&value)?; + } + } + } + ']' => { + self.balances.pop(); + } + _ => {} + } + } + self.cursor = self.buffer.len(); + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + use bytes::Bytes; + use futures_util::stream; + use rand::{thread_rng, Rng}; + + fn split_chunks(text: &str) -> Vec> { + let mut rng = thread_rng(); + let len = text.len(); + let cut1 = rng.gen_range(1..len - 1); + let cut2 = rng.gen_range(cut1 + 1..len); + let chunk1 = text[..cut1].as_bytes().to_vec(); + let chunk2 = text[cut1..cut2].as_bytes().to_vec(); + let chunk3 = text[cut2..].as_bytes().to_vec(); + vec![chunk1, chunk2, chunk3] + } + + macro_rules! assert_json_stream { + ($input:expr, $output:expr) => { + let chunks: Vec<_> = split_chunks($input) + .into_iter() + .map(|chunk| Ok::<_, std::convert::Infallible>(Bytes::from(chunk))) + .collect(); + let stream = stream::iter(chunks); + let mut output = vec![]; + let ret = json_stream(stream, |data| { + output.push(data.to_string()); + Ok(()) + }) + .await; + assert!(ret.is_ok()); + assert_eq!($output.replace("\r\n", "\n"), output.join("\n")) + }; + } + + #[tokio::test] + async fn test_json_stream_ndjson() { + let data = r#"{"key": "value"} +{"key": "value2"} +{"key": "value3"}"#; + assert_json_stream!(data, data); + } + + #[tokio::test] + async fn test_json_stream_array() { + let input = r#"[ +{"key": "value"}, +{"key": "value2"}, +{"key": "value3"},"#; + let output = r#"{"key": "value"} +{"key": "value2"} +{"key": "value3"}"#; + assert_json_stream!(input, output); + } +} -- cgit v1.2.3