diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/stream.rs | 2 | ||||
| -rw-r--r-- | src/serve.rs | 66 |
2 files changed, 45 insertions, 23 deletions
diff --git a/src/client/stream.rs b/src/client/stream.rs index 7e2b5fa..9913f63 100644 --- a/src/client/stream.rs +++ b/src/client/stream.rs @@ -67,7 +67,7 @@ impl SseHandler { } fn safe_ret(&self, ret: Result<()>) -> Result<()> { - if ret.is_err() && self.abort.aborted() { + if self.abort.aborted() { return Ok(()); } ret diff --git a/src/serve.rs b/src/serve.rs index 11b132c..0b0156d 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -14,7 +14,14 @@ use hyper_util::rt::{TokioExecutor, TokioIo}; use parking_lot::RwLock; use serde::Deserialize; use serde_json::{json, Value}; -use std::{convert::Infallible, net::IpAddr, sync::Arc}; +use std::{ + convert::Infallible, + net::IpAddr, + sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }, +}; use tokio::{ net::TcpListener, sync::{ @@ -271,18 +278,18 @@ impl Server { if stream { let (tx, mut rx) = unbounded_channel(); tokio::spawn(async move { - let mut is_first = true; - let (tx2, rx2) = unbounded_channel(); - let mut handler = SseHandler::new(tx2, abort); + let is_first = Arc::new(AtomicBool::new(true)); + let (sse_tx, sse_rx) = unbounded_channel(); + let mut handler = SseHandler::new(sse_tx, abort); async fn map_event( - mut rx: UnboundedReceiver<SseEvent>, + mut sse_rx: UnboundedReceiver<SseEvent>, tx: &UnboundedSender<ResEvent>, - is_first: &mut bool, + is_first: Arc<AtomicBool>, ) { - while let Some(reply_event) = rx.recv().await { - if *is_first { + while let Some(reply_event) = sse_rx.recv().await { + if is_first.load(Ordering::SeqCst) { let _ = tx.send(ResEvent::First(None)); - *is_first = false; + is_first.store(false, Ordering::SeqCst) } match reply_event { SseEvent::Text(text) => { @@ -290,19 +297,41 @@ impl Server { } SseEvent::Done => { let _ = tx.send(ResEvent::Done); + sse_rx.close(); } } } } - tokio::select! { - _ = map_event(rx2, &tx, &mut is_first) => {} - ret = client.chat_completions_streaming_inner(&http_client, &mut handler, data) => { - if let Err(err) = ret { - send_first_event(&tx, Some(format!("{err:?}")), &mut is_first) + async fn chat_completions( + client: &dyn Client, + http_client: &reqwest::Client, + handler: &mut SseHandler, + data: ChatCompletionsData, + tx: &UnboundedSender<ResEvent>, + is_first: Arc<AtomicBool>, + ) { + let ret = client + .chat_completions_streaming_inner(http_client, handler, data) + .await; + if let Err(err) = ret { + if is_first.load(Ordering::SeqCst) { + let _ = tx.send(ResEvent::First(Some(format!("{err:?}")))); + is_first.store(false, Ordering::SeqCst) } - let _ = tx.send(ResEvent::Done); } + let _ = handler.done(); } + tokio::join!( + map_event(sse_rx, &tx, is_first.clone()), + chat_completions( + client.as_ref(), + &http_client, + &mut handler, + data, + &tx, + is_first + ), + ); }); let first_event = rx.recv().await; @@ -446,13 +475,6 @@ enum ResEvent { Done, } -fn send_first_event(tx: &UnboundedSender<ResEvent>, data: Option<String>, is_first: &mut bool) { - if *is_first { - let _ = tx.send(ResEvent::First(data)); - *is_first = false; - } -} - async fn shutdown_signal() { tokio::signal::ctrl_c() .await |
