diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-23 19:28:56 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-23 19:28:56 +0800 |
| commit | 5458150ed3203cf13b0371efa2c791ac696cee93 (patch) | |
| tree | 702d17039d9677246a3b91deb8243e09cd768402 | |
| parent | 2ccbb0f06a4558e15642feb53ba7b2bd72804820 (diff) | |
| download | aichat-5458150ed3203cf13b0371efa2c791ac696cee93.tar.gz | |
fix: json stream parser and refine client modules (#538)
| -rwxr-xr-x | Argcfile.sh | 14 | ||||
| -rw-r--r-- | Cargo.lock | 37 | ||||
| -rw-r--r-- | Cargo.toml | 3 | ||||
| -rw-r--r-- | src/client/claude.rs | 4 | ||||
| -rw-r--r-- | src/client/cloudflare.rs | 4 | ||||
| -rw-r--r-- | src/client/common.rs | 121 | ||||
| -rw-r--r-- | src/client/ernie.rs | 6 | ||||
| -rw-r--r-- | src/client/mod.rs | 4 | ||||
| -rw-r--r-- | src/client/openai.rs | 4 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 14 | ||||
| -rw-r--r-- | src/client/replicate.rs | 8 | ||||
| -rw-r--r-- | src/client/sse_handler.rs | 78 | ||||
| -rw-r--r-- | src/client/stream.rs | 292 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 5 |
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'", @@ -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" @@ -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, |
