summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/common.rs19
-rw-r--r--src/client/ernie.rs8
-rw-r--r--src/client/message.rs69
-rw-r--r--src/client/model.rs12
-rw-r--r--src/client/openai.rs11
-rw-r--r--src/client/palm.rs8
6 files changed, 104 insertions, 23 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index cf5ba9b..2716d87 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -1,7 +1,7 @@
use super::{openai::OpenAIConfig, ClientConfig, Message};
use crate::{
- config::GlobalConfig,
+ config::{GlobalConfig, Input},
render::ReplyHandler,
utils::{
init_tokio_runtime, prompt_input_integer, prompt_input_string, tokenize, AbortSignal,
@@ -50,7 +50,7 @@ macro_rules! register_client {
}
impl $client {
- pub const NAME: &str = $name;
+ pub const NAME: &'static str = $name;
pub fn init(global_config: &$crate::config::GlobalConfig) -> Option<Box<dyn Client>> {
let model = global_config.read().model.clone();
@@ -186,22 +186,22 @@ pub trait Client {
Ok(client)
}
- fn send_message(&self, content: &str) -> Result<String> {
+ 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(content);
+ let content = global_config.read().echo_messages(&input);
return Ok(content);
}
let client = self.build_client()?;
- let data = global_config.read().prepare_send_data(content, false)?;
+ 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, content: &str, handler: &mut ReplyHandler) -> Result<()> {
+ fn send_message_streaming(&self, input: &Input, handler: &mut ReplyHandler) -> Result<()> {
async fn watch_abort(abort: AbortSignal) {
loop {
if abort.aborted() {
@@ -211,12 +211,13 @@ pub trait Client {
}
}
let abort = handler.get_abort();
- init_tokio_runtime()?.block_on(async {
+ 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(content);
+ let content = global_config.read().echo_messages(&input);
let tokens = tokenize(&content);
for token in tokens {
tokio::time::sleep(Duration::from_millis(10)).await;
@@ -225,7 +226,7 @@ pub trait Client {
return Ok(());
}
let client = self.build_client()?;
- let data = global_config.read().prepare_send_data(content, true)?;
+ let data = global_config.read().prepare_send_data(&input, true)?;
self.send_message_streaming_inner(&client, handler, data).await
} => {
handler.done()?;
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 200433c..4bb3435 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -1,4 +1,4 @@
-use super::{ErnieClient, Client, ExtraConfig, PromptType, SendData, Model};
+use super::{ErnieClient, Client, ExtraConfig, PromptType, SendData, Model, MessageContent};
use crate::{
config::GlobalConfig,
@@ -198,8 +198,10 @@ fn build_body(data: SendData, _model: String) -> Value {
if messages[0].role.is_system() {
let system_message = messages.remove(0);
- if let Some(message) = messages.get_mut(0) {
- message.content = format!("{}\n\n{}", system_message.content, message.content)
+ if let (Some(message), MessageContent::Text(system_text)) = (messages.get_mut(0), system_message.content) {
+ if let MessageContent::Text(text) = message.content.clone() {
+ message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text))
+ }
}
}
diff --git a/src/client/message.rs b/src/client/message.rs
index 55b2663..dc8c3e1 100644
--- a/src/client/message.rs
+++ b/src/client/message.rs
@@ -1,16 +1,18 @@
+use crate::config::Input;
+
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct Message {
pub role: MessageRole,
- pub content: String,
+ pub content: MessageContent,
}
impl Message {
- pub fn new(content: &str) -> Self {
+ pub fn new(input: &Input) -> Self {
Self {
role: MessageRole::User,
- content: content.to_string(),
+ content: input.to_message_content(),
}
}
}
@@ -38,6 +40,65 @@ impl MessageRole {
}
}
+#[derive(Debug, Clone, Deserialize, Serialize)]
+#[serde(untagged)]
+pub enum MessageContent {
+ Text(String),
+ Array(Vec<MessageContentPart>),
+}
+
+impl MessageContent {
+ pub fn render_input(&self, resolve_url_fn: impl Fn(&str) -> String) -> String {
+ match self {
+ MessageContent::Text(text) => text.to_string(),
+ MessageContent::Array(list) => {
+ let (mut concated_text, mut files) = (String::new(), vec![]);
+ for item in list {
+ match item {
+ MessageContentPart::Text { text } => {
+ concated_text = format!("{concated_text} {text}")
+ }
+ MessageContentPart::ImageUrl { image_url } => {
+ files.push(resolve_url_fn(&image_url.url))
+ }
+ }
+ }
+ if !concated_text.is_empty() {
+ concated_text = format!(" -- {concated_text}")
+ }
+ format!(".file {}{}", files.join(" "), concated_text)
+ }
+ }
+ }
+
+ pub fn merge_prompt(&mut self, replace_fn: impl Fn(&str) -> String) {
+ match self {
+ MessageContent::Text(text) => *text = replace_fn(text),
+ MessageContent::Array(list) => {
+ if list.is_empty() {
+ list.push(MessageContentPart::Text {
+ text: replace_fn(""),
+ })
+ } else if let Some(MessageContentPart::Text { text }) = list.get_mut(0) {
+ *text = replace_fn(text)
+ }
+ }
+ }
+ }
+}
+
+#[derive(Debug, Clone, Deserialize, Serialize)]
+#[serde(tag = "type", rename_all = "snake_case")]
+pub enum MessageContentPart {
+ Text { text: String },
+ ImageUrl { image_url: ImageUrl },
+}
+
+#[derive(Debug, Clone, Deserialize, Serialize)]
+pub struct ImageUrl {
+ pub url: String,
+}
+
#[cfg(test)]
mod tests {
use super::*;
@@ -45,7 +106,7 @@ mod tests {
#[test]
fn test_serde() {
assert_eq!(
- serde_json::to_string(&Message::new("Hello World")).unwrap(),
+ serde_json::to_string(&Message::new(&Input::from_str("Hello World"))).unwrap(),
"{\"role\":\"user\",\"content\":\"Hello World\"}"
);
}
diff --git a/src/client/model.rs b/src/client/model.rs
index 16fe087..130489d 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -1,4 +1,4 @@
-use super::message::Message;
+use super::message::{Message, MessageContent};
use crate::utils::count_tokens;
@@ -79,7 +79,15 @@ impl Model {
}
pub fn messages_tokens(&self, messages: &[Message]) -> usize {
- messages.iter().map(|v| count_tokens(&v.content)).sum()
+ messages
+ .iter()
+ .map(|v| {
+ match &v.content {
+ MessageContent::Text(text) => count_tokens(text),
+ MessageContent::Array(_) => 0, // TODO
+ }
+ })
+ .sum()
}
pub fn total_tokens(&self, messages: &[Message]) -> usize {
diff --git a/src/client/openai.rs b/src/client/openai.rs
index dbe2661..928c6b3 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -19,13 +19,14 @@ use std::env;
const API_BASE: &str = "https://api.openai.com/v1";
-const MODELS: [(&str, usize); 6] = [
+const MODELS: [(&str, usize); 7] = [
("gpt-3.5-turbo", 4096),
("gpt-3.5-turbo-16k", 16385),
("gpt-3.5-turbo-1106", 16385),
+ ("gpt-4-1106-preview", 128000),
+ ("gpt-4-vision-preview", 128000),
("gpt-4", 8192),
("gpt-4-32k", 32768),
- ("gpt-4-1106-preview", 128000),
];
pub const OPENAI_TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2);
@@ -145,6 +146,12 @@ pub fn openai_build_body(data: SendData, model: String) -> Value {
"model": model,
"messages": messages,
});
+
+ // The default max_tokens of gpt-4-vision-preview is only 16, we need to make it larger
+ if model == "gpt-4-vision-preview" {
+ body["max_tokens"] = json!(4096);
+ }
+
if let Some(v) = temperature {
body["temperature"] = v.into();
}
diff --git a/src/client/palm.rs b/src/client/palm.rs
index a2aec8a..37ed0ae 100644
--- a/src/client/palm.rs
+++ b/src/client/palm.rs
@@ -1,4 +1,4 @@
-use super::{PaLMClient, Client, ExtraConfig, Model, PromptType, SendData, TokensCountFactors, send_message_as_streaming};
+use super::{PaLMClient, Client, ExtraConfig, Model, PromptType, SendData, TokensCountFactors, send_message_as_streaming, MessageContent};
use crate::{config::GlobalConfig, render::ReplyHandler, utils::PromptKind};
@@ -115,8 +115,10 @@ fn build_body(data: SendData, _model: String) -> Value {
if messages[0].role.is_system() {
let system_message = messages.remove(0);
- if let Some(message) = messages.get_mut(0) {
- message.content = format!("{}\n\n{}", system_message.content, message.content)
+ if let (Some(message), MessageContent::Text(system_text)) = (messages.get_mut(0), system_message.content) {
+ if let MessageContent::Text(text) = message.content.clone() {
+ message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text))
+ }
}
}