summaryrefslogtreecommitdiffstats
path: root/src/client/cohere.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-10 20:06:57 +0800
committerGitHub <noreply@github.com>2024-04-10 20:06:57 +0800
commita0bd6e1d5d68c718a0533037364e3df9ad13da96 (patch)
treefb46abab73f9b43c593e75c8f7e2315cd36bce1b /src/client/cohere.rs
parent5915bc2f3a4787cccaa49ba86f670ee740f496fd (diff)
downloadaichat-a0bd6e1d5d68c718a0533037364e3df9ad13da96.tar.gz
refactor: extract json stream handling (#398)
Diffstat (limited to 'src/client/cohere.rs')
-rw-r--r--src/client/cohere.rs60
1 files changed, 9 insertions, 51 deletions
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index ca77035..a93dc71 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -1,13 +1,12 @@
use super::{
- message::*, patch_system_message, Client, CohereClient, ExtraConfig, Model, PromptType,
- SendData, TokensCountFactors,
+ json_stream, message::*, patch_system_message, Client, CohereClient, ExtraConfig, Model,
+ PromptType, SendData, TokensCountFactors,
};
use crate::{render::ReplyHandler, utils::PromptKind};
use anyhow::{bail, Result};
use async_trait::async_trait;
-use futures_util::StreamExt;
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
@@ -107,55 +106,14 @@ pub(crate) async fn send_message_streaming(
let data: Value = res.json().await?;
check_error(&data)?;
} else {
- let mut buffer = vec![];
- let mut cursor = 0;
- let mut start = 0;
- let mut balances = vec![];
- let mut quoting = false;
- let mut stream = res.bytes_stream();
- while let Some(chunk) = stream.next().await {
- let chunk = chunk?;
- let chunk = std::str::from_utf8(&chunk)?;
- buffer.extend(chunk.chars());
- for i in cursor..buffer.len() {
- let ch = buffer[i];
- if quoting {
- if ch == '"' && buffer[i - 1] != '\\' {
- quoting = false;
- }
- continue;
- }
- match ch {
- '"' => quoting = true,
- '{' => {
- if balances.is_empty() {
- start = i;
- }
- balances.push(ch);
- }
- '[' => {
- if start != 0 {
- balances.push(ch);
- }
- }
- '}' => {
- balances.pop();
- if balances.is_empty() {
- let value: String = buffer[start..=i].iter().collect();
- let value: Value = serde_json::from_str(&value)?;
- if let Some("text-generation") = value["event_type"].as_str() {
- handler.text(extract_text(&value)?)?;
- }
- }
- }
- ']' => {
- balances.pop();
- }
- _ => {}
- }
+ let handle = |value: &str| -> Result<()> {
+ let value: Value = serde_json::from_str(value)?;
+ if let Some("text-generation") = value["event_type"].as_str() {
+ handler.text(extract_text(&value)?)?;
}
- cursor = buffer.len();
- }
+ Ok(())
+ };
+ json_stream(res.bytes_stream(), handle).await?;
}
Ok(())
}