diff options
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/claude.rs | 11 | ||||
| -rw-r--r-- | src/client/cohere.rs | 4 | ||||
| -rw-r--r-- | src/client/common.rs | 120 | ||||
| -rw-r--r-- | src/client/ernie.rs | 7 | ||||
| -rw-r--r-- | src/client/gemini.rs | 4 | ||||
| -rw-r--r-- | src/client/mod.rs | 2 | ||||
| -rw-r--r-- | src/client/ollama.rs | 6 | ||||
| -rw-r--r-- | src/client/openai.rs | 4 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 7 | ||||
| -rw-r--r-- | src/client/reply_handler.rs | 65 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 4 |
11 files changed, 164 insertions, 70 deletions
diff --git a/src/client/claude.rs b/src/client/claude.rs index 7a4dd36..4da5128 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -1,11 +1,10 @@ -use super::{patch_system_message, ClaudeClient, Client, ExtraConfig, Model, PromptType, SendData}; - -use crate::{ - client::{ImageUrl, MessageContent, MessageContentPart}, - render::ReplyHandler, - utils::PromptKind, +use super::{ + patch_system_message, ClaudeClient, Client, ExtraConfig, ImageUrl, MessageContent, + MessageContentPart, Model, PromptType, ReplyHandler, SendData, }; +use crate::utils::PromptKind; + use anyhow::{anyhow, bail, Result}; use async_trait::async_trait; use futures_util::StreamExt; diff --git a/src/client/cohere.rs b/src/client/cohere.rs index 6f2e288..a92e238 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,9 +1,9 @@ use super::{ json_stream, message::*, patch_system_message, Client, CohereClient, ExtraConfig, Model, - PromptType, SendData, + PromptType, ReplyHandler, SendData, }; -use crate::{render::ReplyHandler, utils::PromptKind}; +use crate::utils::PromptKind; use anyhow::{bail, Result}; use async_trait::async_trait; diff --git a/src/client/common.rs b/src/client/common.rs index 9171e3a..2206d21 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -1,12 +1,9 @@ -use super::{openai::OpenAIConfig, ClientConfig, Message, MessageContent, Model}; +use super::{openai::OpenAIConfig, ClientConfig, Message, MessageContent, Model, ReplyHandler}; use crate::{ config::{GlobalConfig, Input}, - render::ReplyHandler, - utils::{ - init_tokio_runtime, prompt_input_integer, prompt_input_string, tokenize, AbortSignal, - PromptKind, - }, + render::{render_error, render_stream}, + utils::{prompt_input_integer, prompt_input_string, tokenize, AbortSignal, PromptKind}, }; use anyhow::{Context, Result}; @@ -16,7 +13,7 @@ use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; use std::{env, future::Future, time::Duration}; -use tokio::time::sleep; +use tokio::{sync::mpsc::unbounded_channel, time::sleep}; #[macro_export] macro_rules! register_client { @@ -173,7 +170,7 @@ macro_rules! openai_compatible_client { async fn send_message_streaming_inner( &self, client: &reqwest::Client, - handler: &mut $crate::render::ReplyHandler, + handler: &mut $crate::client::ReplyHandler, data: $crate::client::SendData, ) -> Result<()> { let builder = self.request_builder(client, data)?; @@ -201,7 +198,7 @@ macro_rules! config_get_fn { } #[async_trait] -pub trait Client { +pub trait Client: Sync + Send { fn config(&self) -> (&GlobalConfig, &Option<ExtraConfig>); fn models(&self) -> Vec<Model>; @@ -226,22 +223,24 @@ pub trait Client { Ok(client) } - fn send_message(&self, input: Input) -> Result<String> { - init_tokio_runtime()?.block_on(async { - let global_config = self.config().0; - if global_config.read().dry_run { - let content = global_config.read().echo_messages(&input); - return Ok(content); - } - let client = self.build_client()?; - let data = global_config.read().prepare_send_data(&input, false)?; - self.send_message_inner(&client, data) - .await - .with_context(|| "Failed to get answer") - }) + async fn send_message(&self, input: Input) -> Result<String> { + let global_config = self.config().0; + if global_config.read().dry_run { + let content = global_config.read().echo_messages(&input); + return Ok(content); + } + let client = self.build_client()?; + let data = global_config.read().prepare_send_data(&input, false)?; + self.send_message_inner(&client, data) + .await + .with_context(|| "Failed to get answer") } - fn send_message_streaming(&self, input: &Input, handler: &mut ReplyHandler) -> Result<()> { + async fn send_message_streaming( + &self, + input: &Input, + handler: &mut ReplyHandler, + ) -> Result<()> { async fn watch_abort(abort: AbortSignal) { loop { if abort.aborted() { @@ -252,32 +251,30 @@ pub trait Client { } let abort = handler.get_abort(); let input = input.clone(); - init_tokio_runtime()?.block_on(async move { - tokio::select! { - ret = async { - let global_config = self.config().0; - if global_config.read().dry_run { - let content = global_config.read().echo_messages(&input); - let tokens = tokenize(&content); - for token in tokens { - tokio::time::sleep(Duration::from_millis(10)).await; - handler.text(&token)?; - } - return Ok(()); + tokio::select! { + ret = async { + let global_config = self.config().0; + if global_config.read().dry_run { + let content = global_config.read().echo_messages(&input); + let tokens = tokenize(&content); + for token in tokens { + tokio::time::sleep(Duration::from_millis(10)).await; + handler.text(&token)?; } - let client = self.build_client()?; - let data = global_config.read().prepare_send_data(&input, true)?; - self.send_message_streaming_inner(&client, handler, data).await - } => { - handler.done()?; - ret.with_context(|| "Failed to get answer") + return Ok(()); } - _ = watch_abort(abort.clone()) => { - handler.done()?; - Ok(()) - }, + let client = self.build_client()?; + let data = global_config.read().prepare_send_data(&input, true)?; + self.send_message_streaming_inner(&client, handler, data).await + } => { + handler.done()?; + ret.with_context(|| "Failed to get answer") } - }) + _ = watch_abort(abort.clone()) => { + handler.done()?; + Ok(()) + }, + } } async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String>; @@ -336,6 +333,37 @@ pub fn create_config(list: &[PromptType], client: &str) -> Result<(String, Value Ok((model, clients)) } +pub async fn send_stream( + input: &Input, + client: &dyn Client, + config: &GlobalConfig, + abort: AbortSignal, +) -> Result<String> { + let (tx, rx) = unbounded_channel(); + let mut stream_handler = ReplyHandler::new(tx, abort.clone()); + + let (send_ret, rend_ret) = tokio::join!( + client.send_message_streaming(input, &mut stream_handler), + render_stream(rx, config, abort.clone()), + ); + if let Err(err) = rend_ret { + render_error(err, config.read().highlight); + } + let output = stream_handler.get_buffer().to_string(); + match send_ret { + Ok(_) => { + println!(); + Ok(output) + } + Err(err) => { + if !output.is_empty() { + println!(); + } + Err(err) + } + } +} + #[allow(unused)] pub async fn send_message_as_streaming<F, Fut>( builder: RequestBuilder, diff --git a/src/client/ernie.rs b/src/client/ernie.rs index 060f89b..db6a969 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -1,6 +1,9 @@ -use super::{patch_system_message, Client, ErnieClient, ExtraConfig, Model, PromptType, SendData}; +use super::{ + patch_system_message, Client, ErnieClient, ExtraConfig, Model, PromptType, ReplyHandler, + SendData, +}; -use crate::{render::ReplyHandler, utils::PromptKind}; +use crate::utils::PromptKind; use anyhow::{anyhow, bail, Context, Result}; use async_trait::async_trait; diff --git a/src/client/gemini.rs b/src/client/gemini.rs index bcd10e3..0c60bee 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -1,7 +1,7 @@ use super::vertexai::{build_body, send_message, send_message_streaming}; -use super::{Client, ExtraConfig, GeminiClient, Model, PromptType, SendData}; +use super::{Client, ExtraConfig, GeminiClient, Model, PromptType, ReplyHandler, SendData}; -use crate::{render::ReplyHandler, utils::PromptKind}; +use crate::utils::PromptKind; use anyhow::Result; use async_trait::async_trait; diff --git a/src/client/mod.rs b/src/client/mod.rs index d49ed4e..bd85e74 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -2,10 +2,12 @@ mod common; mod message; mod model; +mod reply_handler; pub use common::*; pub use message::*; pub use model::*; +pub use reply_handler::*; register_client!( (openai, "openai", OpenAIConfig, OpenAIClient), diff --git a/src/client/ollama.rs b/src/client/ollama.rs index 2c51f44..e652634 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -1,9 +1,9 @@ use super::{ message::*, patch_system_message, Client, ExtraConfig, Model, ModelConfig, OllamaClient, - PromptType, SendData, + PromptType, ReplyHandler, SendData, }; -use crate::{render::ReplyHandler, utils::PromptKind}; +use crate::utils::PromptKind; use anyhow::{anyhow, bail, Result}; use async_trait::async_trait; @@ -118,7 +118,7 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHand while let Some(chunk) = stream.next().await { let chunk = chunk?; if chunk.is_empty() { - continue; + continue; } let data: Value = serde_json::from_slice(&chunk)?; if data["done"].is_boolean() { diff --git a/src/client/openai.rs b/src/client/openai.rs index 24c72cc..797ee98 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,6 +1,6 @@ -use super::{ExtraConfig, Model, OpenAIClient, PromptType, SendData}; +use super::{ExtraConfig, Model, OpenAIClient, PromptType, ReplyHandler, SendData}; -use crate::{render::ReplyHandler, utils::PromptKind}; +use crate::utils::PromptKind; use anyhow::{anyhow, bail, Result}; use async_trait::async_trait; diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs index 031abe7..2034736 100644 --- a/src/client/qianwen.rs +++ b/src/client/qianwen.rs @@ -1,11 +1,8 @@ use super::{ - message::*, Client, ExtraConfig, Model, PromptType, QianwenClient, SendData, + message::*, Client, ExtraConfig, Model, PromptType, QianwenClient, ReplyHandler, SendData, }; -use crate::{ - render::ReplyHandler, - utils::{sha256sum, PromptKind}, -}; +use crate::utils::{sha256sum, PromptKind}; use anyhow::{anyhow, bail, Context, Result}; use async_trait::async_trait; diff --git a/src/client/reply_handler.rs b/src/client/reply_handler.rs new file mode 100644 index 0000000..e11ea1d --- /dev/null +++ b/src/client/reply_handler.rs @@ -0,0 +1,65 @@ +use crate::utils::AbortSignal; + +use anyhow::{Context, Result}; +use tokio::sync::mpsc::UnboundedSender; + +pub struct ReplyHandler { + sender: UnboundedSender<ReplyEvent>, + buffer: String, + abort: AbortSignal, +} + +impl ReplyHandler { + pub fn new(sender: UnboundedSender<ReplyEvent>, abort: AbortSignal) -> Self { + Self { + sender, + abort, + buffer: String::new(), + } + } + + pub fn text(&mut self, text: &str) -> Result<()> { + debug!("ReplyText: {}", text); + if text.is_empty() { + return Ok(()); + } + self.buffer.push_str(text); + let ret = self + .sender + .send(ReplyEvent::Text(text.to_string())) + .with_context(|| "Failed to send ReplyEvent:Text"); + self.safe_ret(ret)?; + Ok(()) + } + + pub fn done(&mut self) -> Result<()> { + debug!("ReplyDone"); + let ret = self + .sender + .send(ReplyEvent::Done) + .with_context(|| "Failed to send ReplyEvent::Done"); + self.safe_ret(ret)?; + Ok(()) + } + + pub fn get_buffer(&self) -> &str { + &self.buffer + } + + pub fn get_abort(&self) -> AbortSignal { + self.abort.clone() + } + + fn safe_ret(&self, ret: Result<()>) -> Result<()> { + if ret.is_err() && self.abort.aborted() { + return Ok(()); + } + ret + } +} + +#[derive(Debug)] +pub enum ReplyEvent { + Text(String), + Done, +} diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index ad39c16..88035ec 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -1,9 +1,9 @@ use super::{ json_stream, message::*, patch_system_message, Client, ExtraConfig, Model, PromptType, - SendData, VertexAIClient, + ReplyHandler, SendData, VertexAIClient, }; -use crate::{render::ReplyHandler, utils::PromptKind}; +use crate::utils::PromptKind; use anyhow::{anyhow, bail, Context, Result}; use async_trait::async_trait; |
