summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorProjectMoon <ProjectMoon@users.noreply.github.com>2024-05-29 14:27:07 +0200
committerGitHub <noreply@github.com>2024-05-29 20:27:07 +0800
commit00f3cb182fbf70e99c4ccd21afe8f79359ef3d6d (patch)
treebc56319a647e55ef212a1ac55afcceec1f0043e0
parent4fa92b020aaf0f473a4953b21e14818d6e779208 (diff)
downloadaichat-00f3cb182fbf70e99c4ccd21afe8f79359ef3d6d.tar.gz
refactor: use `json_stream` for ollama to improve reliability (#549)
* Use JSON stream for ollama to improve reliability. Fixes #548. * remove unused import * fix clippy error * format --------- Co-authored-by: sigoden <sigoden@gmail.com>
-rw-r--r--src/client/ollama.rs22
1 files changed, 11 insertions, 11 deletions
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index 37a3667..a650f40 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -1,10 +1,9 @@
use super::{
- catch_error, message::*, Client, CompletionOutput, ExtraConfig, Model, ModelData, ModelPatches,
- OllamaClient, PromptAction, PromptKind, SendData, SseHandler,
+ catch_error, json_stream, message::*, Client, CompletionOutput, ExtraConfig, Model, ModelData,
+ ModelPatches, OllamaClient, PromptAction, PromptKind, SendData, SseHandler,
};
use anyhow::{anyhow, bail, Result};
-use futures_util::StreamExt;
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
@@ -81,14 +80,10 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandle
let data = res.json().await?;
catch_error(&data, status.as_u16())?;
} else {
- let mut stream = res.bytes_stream();
- while let Some(chunk) = stream.next().await {
- let chunk = chunk?;
- if chunk.is_empty() {
- continue;
- }
- let data: Value = serde_json::from_slice(&chunk)?;
+ let handle = |message: &str| -> Result<()> {
+ let data: Value = serde_json::from_str(message)?;
debug!("stream-data: {data}");
+
if data["done"].is_boolean() {
if let Some(text) = data["message"]["content"].as_str() {
handler.text(text)?;
@@ -96,8 +91,13 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandle
} else {
bail!("Invalid response data: {data}")
}
- }
+
+ Ok(())
+ };
+
+ json_stream(res.bytes_stream(), handle).await?;
}
+
Ok(())
}