diff options
| author | ProjectMoon <ProjectMoon@users.noreply.github.com> | 2024-05-29 14:27:07 +0200 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-29 20:27:07 +0800 |
| commit | 00f3cb182fbf70e99c4ccd21afe8f79359ef3d6d (patch) | |
| tree | bc56319a647e55ef212a1ac55afcceec1f0043e0 /src | |
| parent | 4fa92b020aaf0f473a4953b21e14818d6e779208 (diff) | |
| download | aichat-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>
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/ollama.rs | 22 |
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(()) } |
