summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rwxr-xr-xArgcfile.sh14
-rw-r--r--Cargo.lock37
-rw-r--r--Cargo.toml3
-rw-r--r--src/client/claude.rs4
-rw-r--r--src/client/cloudflare.rs4
-rw-r--r--src/client/common.rs121
-rw-r--r--src/client/ernie.rs6
-rw-r--r--src/client/mod.rs4
-rw-r--r--src/client/openai.rs4
-rw-r--r--src/client/qianwen.rs14
-rw-r--r--src/client/replicate.rs8
-rw-r--r--src/client/sse_handler.rs78
-rw-r--r--src/client/stream.rs292
-rw-r--r--src/client/vertexai.rs5
14 files changed, 361 insertions, 233 deletions
diff --git a/Argcfile.sh b/Argcfile.sh
index 88615f9..3b29cda 100755
--- a/Argcfile.sh
+++ b/Argcfile.sh
@@ -217,12 +217,7 @@ chat-cohere() {
-X POST \
-H 'Content-Type: application/json' \
-H "Authorization: Bearer $COHERE_API_KEY" \
---data '{
- "model": "'$argc_model'",
- "message": "'"$*"'",
- "stream": '$stream'
-}
-'
+-d "$(_build_body cohere "$@")"
}
# @cmd List cohere models
@@ -470,6 +465,13 @@ _build_body() {
"stream": '$stream'
}'
;;
+ cohere)
+ echo '{
+ "model": "'$argc_model'",
+ "message": "'"$*"'",
+ "stream": '$stream'
+}'
+ ;;
claude)
echo '{
"model": "'$argc_model'",
diff --git a/Cargo.lock b/Cargo.lock
index 2800501..59fea94 100644
--- a/Cargo.lock
+++ b/Cargo.lock
@@ -61,6 +61,7 @@ dependencies = [
"nu-ansi-term 0.50.0",
"num_cpus",
"parking_lot",
+ "rand",
"reedline",
"reqwest",
"reqwest-eventsource",
@@ -1623,6 +1624,12 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "439ee305def115ba05938db6eb1644ff94165c5ab5e9420d1c1bcedbba909391"
[[package]]
+name = "ppv-lite86"
+version = "0.2.17"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "5b40af805b3121feab8a3c29f04d8ad262fa8e0561883e7653e024ae4479e6de"
+
+[[package]]
name = "proc-macro2"
version = "1.0.82"
source = "registry+https://github.com/rust-lang/crates.io-index"
@@ -1650,6 +1657,36 @@ dependencies = [
]
[[package]]
+name = "rand"
+version = "0.8.5"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "34af8d1a0e25924bc5b7c43c079c942339d8f0a8b57c39049bef581b46327404"
+dependencies = [
+ "libc",
+ "rand_chacha",
+ "rand_core",
+]
+
+[[package]]
+name = "rand_chacha"
+version = "0.3.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "e6c10a63a0fa32252be49d21e7709d4d4baf8d231c2dbce1eaa8141b9b127d88"
+dependencies = [
+ "ppv-lite86",
+ "rand_core",
+]
+
+[[package]]
+name = "rand_core"
+version = "0.6.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "ec0be4795e2f6a28069bec0b5ff3e2ac9bafc99e6a9a7dc3547996c5c816922c"
+dependencies = [
+ "getrandom",
+]
+
+[[package]]
name = "redox_syscall"
version = "0.5.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
diff --git a/Cargo.toml b/Cargo.toml
index f4110c0..d55d81c 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -79,6 +79,9 @@ arboard = { version = "3.3.0", default-features = false, features = ["wayland-da
[target.'cfg(not(any(target_os = "linux", target_os = "android", target_os = "emscripten")))'.dependencies]
arboard = { version = "3.3.0", default-features = false }
+[dev-dependencies]
+rand = "0.8.5"
+
[profile.release]
lto = true
strip = true
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 6296cda..ccedb51 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -1,7 +1,7 @@
use super::{
catch_error, extract_system_message, message::*, sse_stream, ClaudeClient, Client,
CompletionOutput, ExtraConfig, ImageUrl, MessageContent, MessageContentPart, Model, ModelData,
- ModelPatches, PromptAction, PromptKind, SendData, SsMmessage, SseHandler, ToolCall,
+ ModelPatches, PromptAction, PromptKind, SendData, SseHandler, SseMmessage, ToolCall,
};
use anyhow::{bail, Context, Result};
@@ -73,7 +73,7 @@ pub async fn claude_send_message_streaming(
let mut function_name = String::new();
let mut function_arguments = String::new();
let mut function_id = String::new();
- let handle = |message: SsMmessage| -> Result<bool> {
+ let handle = |message: SseMmessage| -> Result<bool> {
let data: Value = serde_json::from_str(&message.data)?;
debug!("stream-data: {data}");
if let Some(typ) = data["type"].as_str() {
diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs
index 3369266..659c7c0 100644
--- a/src/client/cloudflare.rs
+++ b/src/client/cloudflare.rs
@@ -1,6 +1,6 @@
use super::{
catch_error, sse_stream, Client, CloudflareClient, CompletionOutput, ExtraConfig, Model,
- ModelData, ModelPatches, PromptAction, PromptKind, SendData, SsMmessage, SseHandler,
+ ModelData, ModelPatches, PromptAction, PromptKind, SendData, SseHandler, SseMmessage,
};
use anyhow::{anyhow, Result};
@@ -65,7 +65,7 @@ async fn send_message(builder: RequestBuilder) -> Result<CompletionOutput> {
}
async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandler) -> Result<()> {
- let handle = |message: SsMmessage| -> Result<bool> {
+ let handle = |message: SseMmessage| -> Result<bool> {
if message.data == "[DONE]" {
return Ok(true);
}
diff --git a/src/client/common.rs b/src/client/common.rs
index 04844d1..5d16a9b 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -10,11 +10,9 @@ use crate::{
use anyhow::{bail, Context, Result};
use async_trait::async_trait;
use fancy_regex::Regex;
-use futures_util::{Stream, StreamExt};
use indexmap::IndexMap;
use lazy_static::lazy_static;
use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy, RequestBuilder};
-use reqwest_eventsource::{Error as EventSourceError, Event, RequestBuilderExt};
use serde::Deserialize;
use serde_json::{json, Value};
use std::{env, future::Future, time::Duration};
@@ -579,125 +577,6 @@ pub fn maybe_catch_error(data: &Value) -> Result<()> {
Ok(())
}
-#[derive(Debug)]
-pub struct SsMmessage {
- pub event: String,
- pub data: String,
-}
-
-pub async fn sse_stream<F>(builder: RequestBuilder, mut handle: F) -> Result<()>
-where
- F: FnMut(SsMmessage) -> Result<bool>,
-{
- let mut es = builder.eventsource()?;
- while let Some(event) = es.next().await {
- match event {
- Ok(Event::Open) => {}
- Ok(Event::Message(message)) => {
- let message = SsMmessage {
- event: message.event,
- data: message.data,
- };
- if handle(message)? {
- break;
- }
- }
- Err(err) => {
- match err {
- EventSourceError::StreamEnded => {}
- EventSourceError::InvalidStatusCode(status, res) => {
- let text = res.text().await?;
- let data: Value = match text.parse() {
- Ok(data) => data,
- Err(_) => {
- bail!(
- "Invalid response data: {text} (status: {})",
- status.as_u16()
- );
- }
- };
- catch_error(&data, status.as_u16())?;
- }
- EventSourceError::InvalidContentType(header_value, res) => {
- let text = res.text().await?;
- bail!(
- "Invalid response event-stream. content-type: {}, data: {text}",
- header_value.to_str().unwrap_or_default()
- );
- }
- _ => {
- bail!("{}", err);
- }
- }
- es.close();
- }
- }
- }
- Ok(())
-}
-
-pub async fn json_stream<S, F>(mut stream: S, mut handle: F) -> Result<()>
-where
- S: Stream<Item = Result<bytes::Bytes, reqwest::Error>> + Unpin,
- F: FnMut(&str) -> Result<()>,
-{
- let mut buffer = vec![];
- let mut cursor = 0;
- let mut start = 0;
- let mut balances = vec![];
- let mut quoting = false;
- let mut escape = false;
- 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 == '\\' {
- escape = !escape;
- } else {
- if !escape && ch == '"' {
- quoting = false;
- }
- escape = false;
- }
- continue;
- }
- match ch {
- '"' => {
- quoting = true;
- escape = false;
- }
- '{' => {
- 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();
- handle(&value)?;
- }
- }
- ']' => {
- balances.pop();
- }
- _ => {}
- }
- }
- cursor = buffer.len();
- }
- Ok(())
-}
-
fn set_client_config_values(
list: &[PromptAction],
model: &mut String,
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 751eb62..68c4224 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -1,8 +1,8 @@
use super::access_token::*;
use super::{
maybe_catch_error, patch_system_message, sse_stream, Client, CompletionOutput, ErnieClient,
- ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, SendData, SsMmessage,
- SseHandler,
+ ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, SendData, SseHandler,
+ SseMmessage,
};
use anyhow::{anyhow, Context, Result};
@@ -108,7 +108,7 @@ async fn send_message(builder: RequestBuilder) -> Result<CompletionOutput> {
}
async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandler) -> Result<()> {
- let handle = |message: SsMmessage| -> Result<bool> {
+ let handle = |message: SseMmessage| -> Result<bool> {
let data: Value = serde_json::from_str(&message.data)?;
debug!("stream-data: {data}");
if let Some(text) = data["result"].as_str() {
diff --git a/src/client/mod.rs b/src/client/mod.rs
index 0036406..00aa7d0 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -4,14 +4,14 @@ mod access_token;
mod message;
mod model;
mod prompt_format;
-mod sse_handler;
+mod stream;
pub use crate::function::{ToolCall, ToolResults};
pub use crate::utils::PromptKind;
pub use common::*;
pub use message::*;
pub use model::*;
-pub use sse_handler::*;
+pub use stream::*;
register_client!(
(openai, "openai", OpenAIConfig, OpenAIClient),
diff --git a/src/client/openai.rs b/src/client/openai.rs
index c8275cc..881365f 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -1,6 +1,6 @@
use super::{
catch_error, message::*, sse_stream, Client, CompletionOutput, ExtraConfig, Model, ModelData,
- ModelPatches, OpenAIClient, PromptAction, PromptKind, SendData, SsMmessage, SseHandler,
+ ModelPatches, OpenAIClient, PromptAction, PromptKind, SendData, SseHandler, SseMmessage,
ToolCall,
};
@@ -71,7 +71,7 @@ pub async fn openai_send_message_streaming(
let mut function_name = String::new();
let mut function_arguments = String::new();
let mut function_id = String::new();
- let handle = |message: SsMmessage| -> Result<bool> {
+ let handle = |message: SseMmessage| -> Result<bool> {
if message.data == "[DONE]" {
if !function_name.is_empty() {
handler.tool_call(ToolCall::new(
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index ff7f010..2063f20 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -1,7 +1,7 @@
use super::{
maybe_catch_error, message::*, sse_stream, Client, CompletionOutput, ExtraConfig, Model,
- ModelData, ModelPatches, PromptAction, PromptKind, QianwenClient, SendData, SsMmessage,
- SseHandler,
+ ModelData, ModelPatches, PromptAction, PromptKind, QianwenClient, SendData, SseHandler,
+ SseMmessage,
};
use crate::utils::{base64_decode, sha256};
@@ -42,7 +42,7 @@ impl QianwenClient {
let api_key = self.get_api_key()?;
let stream = data.stream;
-
+
let url = match self.model.supports_vision() {
true => API_URL_VL,
false => API_URL,
@@ -64,8 +64,6 @@ impl QianwenClient {
}
}
-
-
#[async_trait]
impl Client for QianwenClient {
client_common_fns!();
@@ -108,14 +106,12 @@ async fn send_message_streaming(
model: &Model,
) -> Result<()> {
let model_name = model.name();
- let handle = |message: SsMmessage| -> Result<bool> {
+ let handle = |message: SseMmessage| -> Result<bool> {
let data: Value = serde_json::from_str(&message.data)?;
maybe_catch_error(&data)?;
debug!("stream-data: {data}");
if model_name == "qwen-long" {
- if let Some(text) =
- data["output"]["choices"][0]["message"]["content"].as_str()
- {
+ if let Some(text) = data["output"]["choices"][0]["message"]["content"].as_str() {
handler.text(text)?;
}
} else if model.supports_vision() {
diff --git a/src/client/replicate.rs b/src/client/replicate.rs
index c549ae8..c39cb07 100644
--- a/src/client/replicate.rs
+++ b/src/client/replicate.rs
@@ -1,7 +1,7 @@
use super::{
- catch_error, prompt_format::*, sse_stream, Client, CompletionOutput, ExtraConfig,
- Model, ModelData, ModelPatches, PromptAction, PromptKind, ReplicateClient, SendData,
- SsMmessage, SseHandler,
+ catch_error, prompt_format::*, sse_stream, Client, CompletionOutput, ExtraConfig, Model,
+ ModelData, ModelPatches, PromptAction, PromptKind, ReplicateClient, SendData, SseHandler,
+ SseMmessage,
};
use anyhow::{anyhow, Result};
@@ -125,7 +125,7 @@ async fn send_message_streaming(
let sse_builder = client.get(stream_url).header("accept", "text/event-stream");
- let handle = |message: SsMmessage| -> Result<bool> {
+ let handle = |message: SseMmessage| -> Result<bool> {
if message.event == "done" {
return Ok(true);
}
diff --git a/src/client/sse_handler.rs b/src/client/sse_handler.rs
deleted file mode 100644
index ddbdbcd..0000000
--- a/src/client/sse_handler.rs
+++ /dev/null
@@ -1,78 +0,0 @@
-use crate::utils::AbortSignal;
-
-use anyhow::{Context, Result};
-use tokio::sync::mpsc::UnboundedSender;
-
-use super::ToolCall;
-
-pub struct SseHandler {
- sender: UnboundedSender<SseEvent>,
- abort: AbortSignal,
- buffer: String,
- tool_calls: Vec<ToolCall>,
-}
-
-impl SseHandler {
- pub fn new(sender: UnboundedSender<SseEvent>, abort: AbortSignal) -> Self {
- Self {
- sender,
- abort,
- buffer: String::new(),
- tool_calls: Vec::new(),
- }
- }
-
- pub fn text(&mut self, text: &str) -> Result<()> {
- // debug!("HandleText: {}", text);
- if text.is_empty() {
- return Ok(());
- }
- self.buffer.push_str(text);
- let ret = self
- .sender
- .send(SseEvent::Text(text.to_string()))
- .with_context(|| "Failed to send ReplyEvent:Text");
- self.safe_ret(ret)?;
- Ok(())
- }
-
- pub fn done(&mut self) -> Result<()> {
- // debug!("HandleDone");
- let ret = self
- .sender
- .send(SseEvent::Done)
- .with_context(|| "Failed to send ReplyEvent::Done");
- self.safe_ret(ret)?;
- Ok(())
- }
-
- pub fn tool_call(&mut self, call: ToolCall) -> Result<()> {
- // debug!("HandleCall: {:?}", call);
- self.tool_calls.push(call);
- Ok(())
- }
-
- pub fn get_abort(&self) -> AbortSignal {
- self.abort.clone()
- }
-
- pub fn take(self) -> (String, Vec<ToolCall>) {
- let Self {
- buffer, tool_calls, ..
- } = self;
- (buffer, tool_calls)
- }
-
- fn safe_ret(&self, ret: Result<()>) -> Result<()> {
- if ret.is_err() && self.abort.aborted() {
- return Ok(());
- }
- ret
- }
-}
-
-#[derive(Debug)]
-pub enum SseEvent {
- Text(String),
- Done,
-}
diff --git a/src/client/stream.rs b/src/client/stream.rs
new file mode 100644
index 0000000..7e2b5fa
--- /dev/null
+++ b/src/client/stream.rs
@@ -0,0 +1,292 @@
+use super::{catch_error, ToolCall};
+use crate::utils::AbortSignal;
+
+use anyhow::{anyhow, bail, Context, Result};
+use futures_util::{Stream, StreamExt};
+use reqwest::RequestBuilder;
+use reqwest_eventsource::{Error as EventSourceError, Event, RequestBuilderExt};
+use serde_json::Value;
+use tokio::sync::mpsc::UnboundedSender;
+
+pub struct SseHandler {
+ sender: UnboundedSender<SseEvent>,
+ abort: AbortSignal,
+ buffer: String,
+ tool_calls: Vec<ToolCall>,
+}
+
+impl SseHandler {
+ pub fn new(sender: UnboundedSender<SseEvent>, abort: AbortSignal) -> Self {
+ Self {
+ sender,
+ abort,
+ buffer: String::new(),
+ tool_calls: Vec::new(),
+ }
+ }
+
+ pub fn text(&mut self, text: &str) -> Result<()> {
+ // debug!("HandleText: {}", text);
+ if text.is_empty() {
+ return Ok(());
+ }
+ self.buffer.push_str(text);
+ let ret = self
+ .sender
+ .send(SseEvent::Text(text.to_string()))
+ .with_context(|| "Failed to send ReplyEvent:Text");
+ self.safe_ret(ret)?;
+ Ok(())
+ }
+
+ pub fn done(&mut self) -> Result<()> {
+ // debug!("HandleDone");
+ let ret = self
+ .sender
+ .send(SseEvent::Done)
+ .with_context(|| "Failed to send ReplyEvent::Done");
+ self.safe_ret(ret)?;
+ Ok(())
+ }
+
+ pub fn tool_call(&mut self, call: ToolCall) -> Result<()> {
+ // debug!("HandleCall: {:?}", call);
+ self.tool_calls.push(call);
+ Ok(())
+ }
+
+ pub fn get_abort(&self) -> AbortSignal {
+ self.abort.clone()
+ }
+
+ pub fn take(self) -> (String, Vec<ToolCall>) {
+ let Self {
+ buffer, tool_calls, ..
+ } = self;
+ (buffer, tool_calls)
+ }
+
+ fn safe_ret(&self, ret: Result<()>) -> Result<()> {
+ if ret.is_err() && self.abort.aborted() {
+ return Ok(());
+ }
+ ret
+ }
+}
+
+#[derive(Debug)]
+pub enum SseEvent {
+ Text(String),
+ Done,
+}
+
+#[derive(Debug)]
+pub struct SseMmessage {
+ pub event: String,
+ pub data: String,
+}
+
+pub async fn sse_stream<F>(builder: RequestBuilder, mut handle: F) -> Result<()>
+where
+ F: FnMut(SseMmessage) -> Result<bool>,
+{
+ let mut es = builder.eventsource()?;
+ while let Some(event) = es.next().await {
+ match event {
+ Ok(Event::Open) => {}
+ Ok(Event::Message(message)) => {
+ let message = SseMmessage {
+ event: message.event,
+ data: message.data,
+ };
+ if handle(message)? {
+ break;
+ }
+ }
+ Err(err) => {
+ match err {
+ EventSourceError::StreamEnded => {}
+ EventSourceError::InvalidStatusCode(status, res) => {
+ let text = res.text().await?;
+ let data: Value = match text.parse() {
+ Ok(data) => data,
+ Err(_) => {
+ bail!(
+ "Invalid response data: {text} (status: {})",
+ status.as_u16()
+ );
+ }
+ };
+ catch_error(&data, status.as_u16())?;
+ }
+ EventSourceError::InvalidContentType(header_value, res) => {
+ let text = res.text().await?;
+ bail!(
+ "Invalid response event-stream. content-type: {}, data: {text}",
+ header_value.to_str().unwrap_or_default()
+ );
+ }
+ _ => {
+ bail!("{}", err);
+ }
+ }
+ es.close();
+ }
+ }
+ }
+ Ok(())
+}
+
+pub async fn json_stream<S, F, E>(mut stream: S, mut handle: F) -> Result<()>
+where
+ S: Stream<Item = Result<bytes::Bytes, E>> + Unpin,
+ F: FnMut(&str) -> Result<()>,
+ E: std::error::Error,
+{
+ let mut parser = JsonStreamParser::default();
+ let mut unparsed_bytes = vec![];
+ while let Some(chunk_bytes) = stream.next().await {
+ let chunk_bytes =
+ chunk_bytes.map_err(|err| anyhow!("Failed to read json stream, {err}"))?;
+ unparsed_bytes.extend(chunk_bytes);
+ match std::str::from_utf8(&unparsed_bytes) {
+ Ok(text) => {
+ parser.process(text, &mut handle)?;
+ unparsed_bytes.clear();
+ }
+ Err(_) => {
+ continue;
+ }
+ }
+ }
+ if !unparsed_bytes.is_empty() {
+ let text = std::str::from_utf8(&unparsed_bytes)?;
+ parser.process(text, &mut handle)?;
+ }
+
+ Ok(())
+}
+
+#[derive(Debug, Default)]
+struct JsonStreamParser {
+ buffer: Vec<char>,
+ cursor: usize,
+ start: Option<usize>,
+ balances: Vec<char>,
+ quoting: bool,
+ escape: bool,
+}
+
+impl JsonStreamParser {
+ fn process<F>(&mut self, text: &str, handle: &mut F) -> Result<()>
+ where
+ F: FnMut(&str) -> Result<()>,
+ {
+ self.buffer.extend(text.chars());
+
+ for i in self.cursor..self.buffer.len() {
+ let ch = self.buffer[i];
+ if self.quoting {
+ if ch == '\\' {
+ self.escape = !self.escape;
+ } else {
+ if !self.escape && ch == '"' {
+ self.quoting = false;
+ }
+ self.escape = false;
+ }
+ continue;
+ }
+ match ch {
+ '"' => {
+ self.quoting = true;
+ self.escape = false;
+ }
+ '{' => {
+ if self.balances.is_empty() {
+ self.start = Some(i);
+ }
+ self.balances.push(ch);
+ }
+ '[' => {
+ if self.start.is_some() {
+ self.balances.push(ch);
+ }
+ }
+ '}' => {
+ self.balances.pop();
+ if self.balances.is_empty() {
+ if let Some(start) = self.start.take() {
+ let value: String = self.buffer[start..=i].iter().collect();
+ handle(&value)?;
+ }
+ }
+ }
+ ']' => {
+ self.balances.pop();
+ }
+ _ => {}
+ }
+ }
+ self.cursor = self.buffer.len();
+ Ok(())
+ }
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ use bytes::Bytes;
+ use futures_util::stream;
+ use rand::{thread_rng, Rng};
+
+ fn split_chunks(text: &str) -> Vec<Vec<u8>> {
+ let mut rng = thread_rng();
+ let len = text.len();
+ let cut1 = rng.gen_range(1..len - 1);
+ let cut2 = rng.gen_range(cut1 + 1..len);
+ let chunk1 = text[..cut1].as_bytes().to_vec();
+ let chunk2 = text[cut1..cut2].as_bytes().to_vec();
+ let chunk3 = text[cut2..].as_bytes().to_vec();
+ vec![chunk1, chunk2, chunk3]
+ }
+
+ macro_rules! assert_json_stream {
+ ($input:expr, $output:expr) => {
+ let chunks: Vec<_> = split_chunks($input)
+ .into_iter()
+ .map(|chunk| Ok::<_, std::convert::Infallible>(Bytes::from(chunk)))
+ .collect();
+ let stream = stream::iter(chunks);
+ let mut output = vec![];
+ let ret = json_stream(stream, |data| {
+ output.push(data.to_string());
+ Ok(())
+ })
+ .await;
+ assert!(ret.is_ok());
+ assert_eq!($output.replace("\r\n", "\n"), output.join("\n"))
+ };
+ }
+
+ #[tokio::test]
+ async fn test_json_stream_ndjson() {
+ let data = r#"{"key": "value"}
+{"key": "value2"}
+{"key": "value3"}"#;
+ assert_json_stream!(data, data);
+ }
+
+ #[tokio::test]
+ async fn test_json_stream_array() {
+ let input = r#"[
+{"key": "value"},
+{"key": "value2"},
+{"key": "value3"},"#;
+ let output = r#"{"key": "value"}
+{"key": "value2"}
+{"key": "value3"}"#;
+ assert_json_stream!(input, output);
+ }
+}
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 67a4b21..1abbffc 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -177,10 +177,7 @@ fn gemini_extract_completion_text(data: &Value) -> Result<CompletionOutput> {
Ok(output)
}
-pub(crate) fn gemini_build_body(
- data: SendData,
- model: &Model,
-) -> Result<Value> {
+pub(crate) fn gemini_build_body(data: SendData, model: &Model) -> Result<Value> {
let SendData {
mut messages,
temperature,