summaryrefslogtreecommitdiffstats
path: root/src/client/vertexai.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/vertexai.rs
parent5915bc2f3a4787cccaa49ba86f670ee740f496fd (diff)
downloadaichat-a0bd6e1d5d68c718a0533037364e3df9ad13da96.tar.gz
refactor: extract json stream handling (#398)
Diffstat (limited to 'src/client/vertexai.rs')
-rw-r--r--src/client/vertexai.rs58
1 files changed, 8 insertions, 50 deletions
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 1c1fd46..babbd23 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -1,6 +1,6 @@
use super::{
- message::*, patch_system_message, Client, ExtraConfig, Model, PromptType, SendData,
- TokensCountFactors, VertexAIClient,
+ json_stream, message::*, patch_system_message, Client, ExtraConfig, Model, PromptType,
+ SendData, TokensCountFactors, VertexAIClient,
};
use crate::{render::ReplyHandler, utils::PromptKind};
@@ -8,7 +8,6 @@ use crate::{render::ReplyHandler, utils::PromptKind};
use anyhow::{anyhow, bail, Context, Result};
use async_trait::async_trait;
use chrono::{Duration, Utc};
-use futures_util::StreamExt;
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
@@ -136,53 +135,12 @@ 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)?;
- handler.text(extract_text(&value)?)?;
- }
- }
- ']' => {
- balances.pop();
- }
- _ => {}
- }
- }
- cursor = buffer.len();
- }
+ let handle = |value: &str| -> Result<()> {
+ let value: Value = serde_json::from_str(value)?;
+ handler.text(extract_text(&value)?)?;
+ Ok(())
+ };
+ json_stream(res.bytes_stream(), handle).await?;
}
Ok(())
}