summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/claude.rs11
-rw-r--r--src/client/cohere.rs4
-rw-r--r--src/client/common.rs120
-rw-r--r--src/client/ernie.rs7
-rw-r--r--src/client/gemini.rs4
-rw-r--r--src/client/mod.rs2
-rw-r--r--src/client/ollama.rs6
-rw-r--r--src/client/openai.rs4
-rw-r--r--src/client/qianwen.rs7
-rw-r--r--src/client/reply_handler.rs65
-rw-r--r--src/client/vertexai.rs4
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;